Get a Hook that controls the restore of the works.
(self)
| 545 | return self._restore_work_queue.dequeue() |
| 546 | |
| 547 | def hook(self): |
| 548 | """Get a Hook that controls the restore of the works.""" |
| 549 | local_work_mgr = self |
| 550 | |
| 551 | class RestoreWorksBarrierHook(training.SessionRunHook): |
| 552 | """Hook that controls the restore of the works |
| 553 | for inference job when failover. |
| 554 | """ |
| 555 | def __init__(self): |
| 556 | super(RestoreWorksBarrierHook, self).__init__() |
| 557 | self._local_work_mgr = local_work_mgr |
| 558 | |
| 559 | def begin(self): |
| 560 | self._update_barrier = \ |
| 561 | self._local_work_mgr.restore_works_barrier.assign(False) |
| 562 | |
| 563 | def after_create_session(self, session, _): |
| 564 | if self._local_work_mgr.restore_work_queue_enqueue: |
| 565 | session.run(self._local_work_mgr.restore_work_queue_enqueue) |
| 566 | old_barrier = session.run(self._local_work_mgr.restore_works_barrier) |
| 567 | new_barrier = session.run(self._update_barrier) |
| 568 | logging.info( |
| 569 | "Update restore_works_barrier:" |
| 570 | "{}->{}".format(old_barrier, new_barrier)) |
| 571 | |
| 572 | def end(self, session): |
| 573 | session.run(self._local_work_mgr.close_restore_work_queue) |
| 574 | logging.info("Close restore_work_queue.") |
| 575 | |
| 576 | return RestoreWorksBarrierHook() |
| 577 | |
| 578 | def local_workqueue_empty(self): |
| 579 | """Return the truth value if local workqueue empty.""" |