improve serialization and add DistArray class

This commit is contained in:
Robert Nishihara
2016-03-16 18:11:43 -07:00
parent 2e48ec0a70
commit 2500cbaf72
9 changed files with 170 additions and 34 deletions
+44 -2
View File
@@ -1,13 +1,14 @@
import unittest
import orchpy
import orchpy.serialization as serialization
import orchpy.services as services
import orchpy.worker as worker
import numpy as np
import time
import subprocess32 as subprocess
import os
import arrays.single as single
import arrays.dist as dist
from google.protobuf.text_format import *
@@ -63,7 +64,7 @@ class ArraysSingleTest(unittest.TestCase):
time.sleep(0.2)
worker.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker1_port))
orchpy.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker1_port))
test_dir = os.path.dirname(os.path.abspath(__file__))
test_path = os.path.join(test_dir, "testrecv.py")
@@ -98,5 +99,46 @@ class ArraysSingleTest(unittest.TestCase):
services.cleanup()
class ArraysDistTest(unittest.TestCase):
def testMethods(self):
x = dist.DistArray()
x.construct([2, 3, 4], float, np.array([[[orchpy.lib.ObjRef(0)]]]))
capsule = serialization.serialize(x)
y = serialization.deserialize(capsule)
self.assertEqual(x.shape, y.shape)
self.assertEqual(x.dtype, y.dtype)
self.assertEqual(x.objrefs[0, 0, 0].val, y.objrefs[0, 0, 0].val)
def testAssemble(self):
scheduler_port = new_scheduler_port()
objstore_port = new_objstore_port()
worker1_port = new_worker_port()
worker2_port = new_worker_port()
services.start_scheduler(address(IP_ADDRESS, scheduler_port))
time.sleep(0.1)
services.start_objstore(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port))
time.sleep(0.2)
orchpy.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker1_port))
test_dir = os.path.dirname(os.path.abspath(__file__))
test_path = os.path.join(test_dir, "testrecv.py")
services.start_worker(test_path, address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker2_port))
time.sleep(0.2)
a = single.ones([dist.BLOCK_SIZE, dist.BLOCK_SIZE])
b = single.zeros([dist.BLOCK_SIZE, dist.BLOCK_SIZE])
x = dist.DistArray()
x.construct([2 * dist.BLOCK_SIZE, dist.BLOCK_SIZE], float, np.array([[a], [b]]))
self.assertTrue(np.alltrue(x.assemble() == np.vstack([np.ones([dist.BLOCK_SIZE, dist.BLOCK_SIZE]), np.zeros([dist.BLOCK_SIZE, dist.BLOCK_SIZE])])))
services.cleanup()
if __name__ == '__main__':
unittest.main()
+27 -26
View File
@@ -1,5 +1,6 @@
import unittest
import orchpy
import orchpy.serialization as serialization
import orchpy.services as services
import orchpy.worker as worker
import numpy as np
@@ -50,14 +51,14 @@ def new_objstore_port():
class SerializationTest(unittest.TestCase):
def roundTripTest(self, data):
serialized = orchpy.lib.serialize_object(data)
result = orchpy.lib.deserialize_object(serialized)
serialized = serialization.serialize(data)
result = serialization.deserialize(serialized)
self.assertEqual(data, result)
def numpyTypeTest(self, typ):
a = np.random.randint(0, 10, size=(100, 100)).astype(typ)
b = orchpy.lib.serialize_object(a)
c = orchpy.lib.deserialize_object(b)
b = serialization.serialize(a)
c = serialization.deserialize(b)
self.assertTrue((a == c).all())
def testSerialize(self):
@@ -68,8 +69,8 @@ class SerializationTest(unittest.TestCase):
self.roundTripTest((1.0, "hi"))
a = np.zeros((100, 100))
res = orchpy.lib.serialize_object(a)
b = orchpy.lib.deserialize_object(res)
res = serialization.serialize(a)
b = serialization.deserialize(res)
self.assertTrue((a == b).all())
self.numpyTypeTest('int8')
@@ -80,8 +81,8 @@ class SerializationTest(unittest.TestCase):
self.numpyTypeTest('float64')
a = np.array([[orchpy.lib.ObjRef(0), orchpy.lib.ObjRef(1)], [orchpy.lib.ObjRef(41), orchpy.lib.ObjRef(42)]])
capsule = orchpy.lib.serialize_object(a)
result = orchpy.lib.deserialize_object(capsule)
capsule = serialization.serialize(a)
result = serialization.deserialize(capsule)
self.assertTrue((a == result).all())
class OrchPyLibTest(unittest.TestCase):
@@ -101,7 +102,7 @@ class OrchPyLibTest(unittest.TestCase):
w = worker.Worker()
worker.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker_port), w)
orchpy.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker_port), w)
w.put_object(orchpy.lib.ObjRef(0), 'hello world')
result = w.get_object(orchpy.lib.ObjRef(0))
@@ -134,22 +135,22 @@ class ObjStoreTest(unittest.TestCase):
objstore2_stub = connect_to_objstore(IP_ADDRESS, objstore2_port)
worker1 = worker.Worker()
worker.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore1_port), address(IP_ADDRESS, worker1_port), worker1)
orchpy.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore1_port), address(IP_ADDRESS, worker1_port), worker1)
worker2 = worker.Worker()
worker.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore2_port), address(IP_ADDRESS, worker2_port), worker2)
orchpy.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore2_port), address(IP_ADDRESS, worker2_port), worker2)
# pushing and pulling an object shouldn't change it
for data in ["h", "h" * 10000, 0, 0.0]:
objref = worker.push(data, worker1)
result = worker.pull(objref, worker1)
objref = orchpy.push(data, worker1)
result = orchpy.pull(objref, worker1)
self.assertEqual(result, data)
# pushing an object, shipping it to another worker, and pulling it shouldn't change it
for data in ["h", "h" * 10000, 0, 0.0]:
objref = worker.push(data, worker1)
objref = orchpy.push(data, worker1)
response = objstore1_stub.DeliverObj(orchestra_pb2.DeliverObjRequest(objref=objref.val, objstore_address=address(IP_ADDRESS, objstore2_port)), TIMEOUT_SECONDS)
result = worker.pull(objref, worker2)
result = orchpy.pull(objref, worker2)
self.assertEqual(result, data)
services.cleanup()
@@ -176,7 +177,7 @@ class SchedulerTest(unittest.TestCase):
time.sleep(0.2)
worker1 = worker.Worker()
worker.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker1_port), worker1)
orchpy.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker1_port), worker1)
test_dir = os.path.dirname(os.path.abspath(__file__))
test_path = os.path.join(test_dir, "testrecv.py")
@@ -189,7 +190,7 @@ class SchedulerTest(unittest.TestCase):
time.sleep(0.2)
value_after = worker.pull(objref[0], worker1)
value_after = orchpy.pull(objref[0], worker1)
self.assertEqual(value_before, value_after)
time.sleep(0.1)
@@ -214,30 +215,30 @@ class WorkerTest(unittest.TestCase):
time.sleep(0.2)
worker1 = worker.Worker()
worker.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker1_port), worker1)
orchpy.connect(address(IP_ADDRESS, scheduler_port), address(IP_ADDRESS, objstore_port), address(IP_ADDRESS, worker1_port), worker1)
for i in range(100):
value_before = i * 10 ** 6
objref = worker.push(value_before, worker1)
value_after = worker.pull(objref, worker1)
objref = orchpy.push(value_before, worker1)
value_after = orchpy.pull(objref, worker1)
self.assertEqual(value_before, value_after)
for i in range(100):
value_before = i * 10 ** 6 * 1.0
objref = worker.push(value_before, worker1)
value_after = worker.pull(objref, worker1)
objref = orchpy.push(value_before, worker1)
value_after = orchpy.pull(objref, worker1)
self.assertEqual(value_before, value_after)
for i in range(100):
value_before = "h" * i
objref = worker.push(value_before, worker1)
value_after = worker.pull(objref, worker1)
objref = orchpy.push(value_before, worker1)
value_after = orchpy.pull(objref, worker1)
self.assertEqual(value_before, value_after)
for i in range(100):
value_before = [1] * i
objref = worker.push(value_before, worker1)
value_after = worker.pull(objref, worker1)
objref = orchpy.push(value_before, worker1)
value_after = orchpy.pull(objref, worker1)
self.assertEqual(value_before, value_after)
services.cleanup()