diff --git a/lib/orchpy/orchpy/services.py b/lib/orchpy/orchpy/services.py index faf431283..cb6cdaca7 100644 --- a/lib/orchpy/orchpy/services.py +++ b/lib/orchpy/orchpy/services.py @@ -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) diff --git a/test/arrays_test.py b/test/arrays_test.py index dd442be3d..4a2644694 100644 --- a/test/arrays_test.py +++ b/test/arrays_test.py @@ -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) diff --git a/test/runtest.py b/test/runtest.py index aace05526..858fd5e28 100644 --- a/test/runtest.py +++ b/test/runtest.py @@ -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)