| 751 | } |
| 752 | |
| 753 | void ltfb::on_batch_begin(model* m) |
| 754 | { |
| 755 | auto& local_model = *m; |
| 756 | auto& context = local_model.get_execution_context(); |
| 757 | auto&& comm = *local_model.get_comm(); |
| 758 | |
| 759 | // Check whether to start LTFB round |
| 760 | const auto mode = context.get_execution_mode(); |
| 761 | const auto step = context.get_step(); |
| 762 | if (mode != execution_mode::training || step == 0) { |
| 763 | return; |
| 764 | } |
| 765 | |
| 766 | // Print message |
| 767 | const auto message_prefix = |
| 768 | (std::string{} + "LTFB (" + "model \"" + local_model.get_name() + "\", " + |
| 769 | "step " + std::to_string(step) + "): "); |
| 770 | if (comm.am_world_master()) { |
| 771 | std::cout << message_prefix + "starting tournament...\n"; |
| 772 | } |
| 773 | |
| 774 | // Determine partner model for tournament |
| 775 | const El::Int local_trainer = comm.get_trainer_rank(); |
| 776 | const El::Int partner_trainer = get_partner_trainer(comm, message_prefix); |
| 777 | |
| 778 | // Evaluate local model |
| 779 | if (comm.am_world_master()) { |
| 780 | std::cout << message_prefix + "evaluating local model...\n"; |
| 781 | } |
| 782 | auto local_score = evaluate(local_model, m_metric_name); |
| 783 | |
| 784 | // Get model from partner trainer |
| 785 | if (comm.am_world_master()) { |
| 786 | std::cout << message_prefix + "exchanging model data...\n"; |
| 787 | } |
| 788 | |
| 789 | model partner_model(local_model); |
| 790 | if (comm_algo_) |
| 791 | comm_algo_->exchange_models(partner_model, partner_trainer, step); |
| 792 | else |
| 793 | LBANN_ERROR("No communication algorithm."); |
| 794 | |
| 795 | // Evaluate partner model |
| 796 | if (comm.am_world_master()) { |
| 797 | std::cout << message_prefix + "evaluating partner model...\n"; |
| 798 | } |
| 799 | auto partner_score = evaluate(partner_model, m_metric_name); |
| 800 | |
| 801 | // Choose tournament winner |
| 802 | // Note: restore local model data if it got a better score. |
| 803 | El::Int tournament_winner = local_trainer; |
| 804 | if ((m_low_score_wins && partner_score <= local_score) || |
| 805 | (!m_low_score_wins && partner_score >= local_score) || |
| 806 | (!std::isfinite(local_score) && std::isfinite(partner_score))) { |
| 807 | tournament_winner = partner_trainer; |
| 808 | |
| 809 | /// @todo Use move assignment operator once LTFB is moved into a |
| 810 | /// training algorithm |
no test coverage detected