r"""Distributed Server Method. Used for exchange information between distributed nodes. Args: mm_server_port: multiple machine rpc server port.
| 13 | |
| 14 | |
| 15 | class Methods: |
| 16 | r"""Distributed Server Method. |
| 17 | Used for exchange information between distributed nodes. |
| 18 | |
| 19 | Args: |
| 20 | mm_server_port: multiple machine rpc server port. |
| 21 | """ |
| 22 | |
| 23 | def __init__(self, mm_server_port): |
| 24 | self.lock = threading.Lock() |
| 25 | self.mm_server_port = mm_server_port |
| 26 | self.dict_is_grad = defaultdict(partial(Future, True)) |
| 27 | self.dict_remote_tracer = defaultdict(partial(Future, True)) |
| 28 | self.dict_pack_list = defaultdict(partial(Future, False)) |
| 29 | self.dict_barrier_counter = defaultdict(int) |
| 30 | self.dict_barrier_event = defaultdict(threading.Event) |
| 31 | self.user_dict = defaultdict(partial(Future, False)) |
| 32 | self.bcast_dict = {} |
| 33 | |
| 34 | def connect(self): |
| 35 | r"""Method for checking connection success.""" |
| 36 | return True |
| 37 | |
| 38 | def get_mm_server_port(self): |
| 39 | r"""Get multiple machine rpc server port.""" |
| 40 | return self.mm_server_port |
| 41 | |
| 42 | def set_is_grad(self, key, is_grad): |
| 43 | r"""Mark send/recv need gradiants by key. |
| 44 | |
| 45 | Args: |
| 46 | key: key to match send/recv op. |
| 47 | is_grad: whether this op need grad. |
| 48 | """ |
| 49 | with self.lock: |
| 50 | future = self.dict_is_grad[key] |
| 51 | future.set(is_grad) |
| 52 | return True |
| 53 | |
| 54 | def check_is_grad(self, key): |
| 55 | r"""Check whether send/recv need gradiants. |
| 56 | |
| 57 | Args: |
| 58 | key: key to match send/recv op. |
| 59 | """ |
| 60 | with self.lock: |
| 61 | future = self.dict_is_grad[key] |
| 62 | ret = future.get() |
| 63 | with self.lock: |
| 64 | del self.dict_is_grad[key] |
| 65 | return ret |
| 66 | |
| 67 | def set_remote_tracer(self, key, tracer_set): |
| 68 | r"""Set tracer dict for tracing send/recv op. |
| 69 | |
| 70 | Args: |
| 71 | key: key to match send/recv op. |
| 72 | tracer_set: valid tracer set. |
no outgoing calls
no test coverage detected