Skip to content

Instantly share code, notes, and snippets.

@kuenishi
Last active August 21, 2019 14:41
Show Gist options
  • Select an option

  • Save kuenishi/f23cd0d0ce5a961fe2dda66bb44e654d to your computer and use it in GitHub Desktop.

Select an option

Save kuenishi/f23cd0d0ce5a961fe2dda66bb44e654d to your computer and use it in GitHub Desktop.
aa
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