mirror of
https://github.com/wassname/ray.git
synced 2026-08-02 13:01:01 +08:00
change services to start variable number of object stores (#55)
This commit is contained in:
committed by
Philipp Moritz
parent
5c014a9857
commit
a12e6fd373
@@ -9,7 +9,7 @@ import orchpy.worker as worker
|
||||
_services_path = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
all_processes = []
|
||||
driver = None
|
||||
drivers = []
|
||||
|
||||
IP_ADDRESS = "127.0.0.1"
|
||||
TIMEOUT_SECONDS = 5
|
||||
@@ -56,12 +56,12 @@ def cleanup():
|
||||
print "Termination attempt failed, giving up."
|
||||
all_processes = []
|
||||
|
||||
global driver
|
||||
if driver is not None:
|
||||
global drivers
|
||||
for driver in drivers:
|
||||
orchpy.disconnect(driver)
|
||||
else:
|
||||
if len(drivers) == 0:
|
||||
orchpy.disconnect()
|
||||
driver = None
|
||||
drivers = []
|
||||
|
||||
# atexit.register(cleanup)
|
||||
|
||||
@@ -81,21 +81,35 @@ def start_worker(test_path, scheduler_address, objstore_address, worker_address)
|
||||
"--worker-address=" + worker_address])
|
||||
all_processes.append((p, worker_address))
|
||||
|
||||
def start_cluster(driver_worker=None, num_workers=0, worker_path=None):
|
||||
global driver
|
||||
if num_workers > 0 and worker_path is None:
|
||||
raise Exception("Attempting to start a cluster with some workers, but `worker_path` is None.")
|
||||
def start_cluster(return_drivers=False, num_objstores=1, num_workers_per_objstore=0, worker_path=None):
|
||||
global drivers
|
||||
if num_workers_per_objstore > 0 and worker_path is None:
|
||||
raise Exception("Attempting to start a cluster with {} workers per object store, but `worker_path` is None.".format(num_workers_per_objstore))
|
||||
if num_workers_per_objstore > 0 and num_objstores < 1:
|
||||
raise Exception("Attempting to start a cluster with {} workers per object store, but `num_objstores` is {}.".format(num_objstores))
|
||||
scheduler_address = address(IP_ADDRESS, new_scheduler_port())
|
||||
objstore_address = address(IP_ADDRESS, new_objstore_port())
|
||||
start_scheduler(scheduler_address)
|
||||
time.sleep(0.1)
|
||||
start_objstore(scheduler_address, objstore_address)
|
||||
time.sleep(0.2)
|
||||
if driver_worker is not None:
|
||||
orchpy.connect(scheduler_address, objstore_address, address(IP_ADDRESS, new_worker_port()), driver_worker)
|
||||
driver = driver_worker
|
||||
objstore_addresses = []
|
||||
# create objstores
|
||||
for i in range(num_objstores):
|
||||
objstore_address = address(IP_ADDRESS, new_objstore_port())
|
||||
objstore_addresses.append(objstore_address)
|
||||
start_objstore(scheduler_address, objstore_address)
|
||||
time.sleep(0.2)
|
||||
for _ in range(num_workers_per_objstore):
|
||||
start_worker(worker_path, scheduler_address, objstore_address, address(IP_ADDRESS, new_worker_port()))
|
||||
time.sleep(0.3)
|
||||
# create drivers
|
||||
if return_drivers:
|
||||
driver_workers = []
|
||||
for i in range(num_objstores):
|
||||
driver_worker = worker.Worker()
|
||||
orchpy.connect(scheduler_address, objstore_address, address(IP_ADDRESS, new_worker_port()), driver_worker)
|
||||
driver_workers.append(driver_worker)
|
||||
drivers.append(driver_worker)
|
||||
time.sleep(0.5)
|
||||
return driver_workers
|
||||
else:
|
||||
orchpy.connect(scheduler_address, objstore_address, address(IP_ADDRESS, new_worker_port()))
|
||||
for _ in range(num_workers):
|
||||
start_worker(worker_path, scheduler_address, objstore_address, address(IP_ADDRESS, new_worker_port()))
|
||||
time.sleep(0.5)
|
||||
orchpy.connect(scheduler_address, objstore_addresses[0], address(IP_ADDRESS, new_worker_port()))
|
||||
time.sleep(0.5)
|
||||
|
||||
+4
-5
@@ -22,7 +22,7 @@ class ArraysSingleTest(unittest.TestCase):
|
||||
def testMethods(self):
|
||||
test_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
test_path = os.path.join(test_dir, "testrecv.py")
|
||||
services.start_cluster(num_workers=1, worker_path=test_path)
|
||||
services.start_cluster(return_drivers=False, num_workers_per_objstore=1, worker_path=test_path)
|
||||
|
||||
# test eye
|
||||
ref = single.eye(3, "float")
|
||||
@@ -54,8 +54,7 @@ class ArraysSingleTest(unittest.TestCase):
|
||||
class ArraysDistTest(unittest.TestCase):
|
||||
|
||||
def testSerialization(self):
|
||||
w = worker.Worker()
|
||||
services.start_cluster(driver_worker=w)
|
||||
[w] = services.start_cluster(return_drivers=True)
|
||||
|
||||
x = dist.DistArray()
|
||||
x.construct([2, 3, 4], np.array([[[orchpy.push(0, w)]]]))
|
||||
@@ -69,7 +68,7 @@ class ArraysDistTest(unittest.TestCase):
|
||||
def testAssemble(self):
|
||||
test_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
test_path = os.path.join(test_dir, "testrecv.py")
|
||||
services.start_cluster(num_workers=1, worker_path=test_path)
|
||||
services.start_cluster(return_drivers=False, num_workers_per_objstore=1, worker_path=test_path)
|
||||
|
||||
a = single.ones([dist.BLOCK_SIZE, dist.BLOCK_SIZE], "float")
|
||||
b = single.zeros([dist.BLOCK_SIZE, dist.BLOCK_SIZE], "float")
|
||||
@@ -82,7 +81,7 @@ class ArraysDistTest(unittest.TestCase):
|
||||
def testMethods(self):
|
||||
test_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
test_path = os.path.join(test_dir, "testrecv.py")
|
||||
services.start_cluster(num_workers=8, worker_path=test_path)
|
||||
services.start_cluster(return_drivers=False, num_workers_per_objstore=8, worker_path=test_path)
|
||||
|
||||
x = dist.zeros([9, 25, 51], "float")
|
||||
y = dist.assemble(x)
|
||||
|
||||
+8
-13
@@ -31,8 +31,7 @@ class SerializationTest(unittest.TestCase):
|
||||
self.assertTrue((a == c).all())
|
||||
|
||||
def testSerialize(self):
|
||||
w = worker.Worker()
|
||||
services.start_cluster(driver_worker=w)
|
||||
[w] = services.start_cluster(return_drivers=True)
|
||||
|
||||
self.roundTripTest(w, [1, "hello", 3.0])
|
||||
self.roundTripTest(w, 42)
|
||||
@@ -70,13 +69,12 @@ class ObjStoreTest(unittest.TestCase):
|
||||
|
||||
# Test setting up object stores, transfering data between them and retrieving data to a client
|
||||
def testObjStore(self):
|
||||
w = worker.Worker()
|
||||
services.start_cluster(driver_worker=w)
|
||||
[w1, w2] = services.start_cluster(return_drivers=True, num_objstores=2, num_workers_per_objstore=0)
|
||||
|
||||
# pushing and pulling an object shouldn't change it
|
||||
for data in ["h", "h" * 10000, 0, 0.0]:
|
||||
objref = orchpy.push(data, w)
|
||||
result = orchpy.pull(objref, w)
|
||||
objref = orchpy.push(data, w1)
|
||||
result = orchpy.pull(objref, w1)
|
||||
self.assertEqual(result, data)
|
||||
|
||||
# pushing an object, shipping it to another worker, and pulling it shouldn't change it
|
||||
@@ -93,8 +91,7 @@ class SchedulerTest(unittest.TestCase):
|
||||
def testCall(self):
|
||||
test_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
test_path = os.path.join(test_dir, "testrecv.py")
|
||||
w = worker.Worker()
|
||||
services.start_cluster(driver_worker=w, num_workers=1, worker_path=test_path)
|
||||
[w] = services.start_cluster(return_drivers=True, num_workers_per_objstore=1, worker_path=test_path)
|
||||
|
||||
value_before = "test_string"
|
||||
objref = w.remote_call("test_functions.print_string", [value_before])
|
||||
@@ -111,8 +108,7 @@ class SchedulerTest(unittest.TestCase):
|
||||
class WorkerTest(unittest.TestCase):
|
||||
|
||||
def testPushPull(self):
|
||||
w = worker.Worker()
|
||||
services.start_cluster(driver_worker=w)
|
||||
[w] = services.start_cluster(return_drivers=True)
|
||||
|
||||
for i in range(100):
|
||||
value_before = i * 10 ** 6
|
||||
@@ -143,10 +139,9 @@ class WorkerTest(unittest.TestCase):
|
||||
class APITest(unittest.TestCase):
|
||||
|
||||
def testObjRefAliasing(self):
|
||||
w = worker.Worker()
|
||||
test_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
test_path = os.path.join(test_dir, "testrecv.py")
|
||||
services.start_cluster(num_workers=3, worker_path=test_path, driver_worker=w)
|
||||
[w] = services.start_cluster(return_drivers=True, num_workers_per_objstore=3, worker_path=test_path)
|
||||
|
||||
objref = w.remote_call("test_functions.test_alias_f", [])
|
||||
self.assertTrue(np.alltrue(orchpy.pull(objref[0], w) == np.ones([3, 4, 5])))
|
||||
@@ -162,7 +157,7 @@ class ReferenceCountingTest(unittest.TestCase):
|
||||
def testDeallocation(self):
|
||||
test_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
test_path = os.path.join(test_dir, "testrecv.py")
|
||||
services.start_cluster(num_workers=3, worker_path=test_path)
|
||||
services.start_cluster(return_drivers=False, num_workers_per_objstore=3, worker_path=test_path)
|
||||
|
||||
x = test_functions.test_alias_f()
|
||||
orchpy.pull(x)
|
||||
|
||||
Reference in New Issue
Block a user