Last active
August 21, 2019 14:41
-
-
Save kuenishi/f23cd0d0ce5a961fe2dda66bb44e654d to your computer and use it in GitHub Desktop.
aa
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| from chainer.training import extension | |
| from chainer.training import extensions | |
| class ReplicaSets: | |
| def __init__(rank, replica_sets): | |
| self.rank = rank | |
| self.replica_sets = replica_sets | |
| self.master = None | |
| self._replicas = [] | |
| for replica_set in replica_sets: | |
| # if replica_set == 'remain': | |
| # replace with remaining set | |
| if self.rank in replica_set: | |
| self.master = replica_set[0] | |
| self._replicas = replica_set | |
| break | |
| @property | |
| def is_master(self): | |
| return self.master == self.rank | |
| def replicas(self): | |
| return self._replicas | |
| def _MultiNodeSnapshot(extension.Extension): | |
| def __init__(self, comm, snapshot, replica_sets, parallel_read=False): | |
| self.comm = comm | |
| self.snapshot = snapshot | |
| self.rs = ReplicaSets(comm.rank, replica_sets) | |
| self.parallel_read = parallel_read | |
| # TODO; Override snapshot filename generation technique | |
| # esp. when it's callable | |
| filename = snapshot.filename + ('.{}'.format(self.rank)) | |
| snapshot.filename = filename | |
| def initialize(self, trainer): | |
| if self.rs.is_master: | |
| self.snapshot.initialize(trainer) | |
| target = trainer if self._target is None else self._target | |
| # Broadcast the target here | |
| if rs.is_master: | |
| # if parallel_read do parallel read, but sharing just names | |
| buf = io.BytesIO() | |
| npz.save_npz(buf, target) | |
| for rank in rs.replicas(): | |
| if rank == comm.rank: | |
| continue | |
| comm.send_obj(buf.buffer(), rank) | |
| else: | |
| buf = comm.recv_obj(rs.master) | |
| npz.load_npz(buf, target) | |
| def on_error(self, trainer, e, t): | |
| if self.rs.is_master: | |
| self.snapshot.on_error(self, trainer, e, t) | |
| def __call__(self, trainer): | |
| if self.rs.is_master: | |
| self.snapshot(trainer) | |
| def finalize(self): | |
| if self.rs.is_master: | |
| self.snapshot.finalize() | |
| snapshot = extensions.snapshot(target, ...) | |
| snapshot = _MultiNodeSnapshot(comm, snapshot, replica_sets = [0, 'remain']) | |
| trainer.extend(snapshot) | |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment