| 638 | } |
| 639 | |
| 640 | bool LoopGrad::GradDesc::bind_grad_var(VarNode* owner_wrt, VarNode* owner_dest) { |
| 641 | auto&& input_ogvar2info = m_fwd_graph_modifier.input_ogvar2info(); |
| 642 | auto info_iter = input_ogvar2info.find(owner_wrt); |
| 643 | if (info_iter == input_ogvar2info.end()) { |
| 644 | // caused by input vars not needed by grad |
| 645 | return false; |
| 646 | } |
| 647 | auto&& info = m_fwd_graph_modifier.input_ogvar2info().at(owner_wrt); |
| 648 | auto&& assignee2info = m_fwd_graph_modifier.assignee2info(); |
| 649 | bool nonzero = false; |
| 650 | for (auto i : info.subgraph_var) { |
| 651 | auto grad = cg::grad(m_grad_virtual_loss, i, false, false); |
| 652 | if (!grad.node()) |
| 653 | continue; |
| 654 | nonzero = true; |
| 655 | bool sum_last; |
| 656 | if (!assignee2info.count(i)) { |
| 657 | // sum all intermediate grads |
| 658 | mgb_assert(i->owner_opr()->same_type<InputMaker>()); |
| 659 | sum_last = false; |
| 660 | } else { |
| 661 | mgb_assert(!i->owner_opr()->same_type<InputMaker>()); |
| 662 | grad.node()->add_flag(VarNode::Flag::NO_MEM_RECLAIM); |
| 663 | sum_last = true; |
| 664 | } |
| 665 | auto grad_sum_recorder = std::make_unique<OutputRecorderSumIntoDest>( |
| 666 | sum_last, &info.grad_dest_summed, owner_dest); |
| 667 | grad = grad_sum_recorder->optimize_grad_var(grad); |
| 668 | do_add_output(grad, std::move(grad_sum_recorder)); |
| 669 | mgb_assert(m_output_record_spec_no_dedup.back()->var_sub() == grad.node()); |
| 670 | const_cast<OutputRecordSpecItem&>(*m_output_record_spec_no_dedup.back()) |
| 671 | .bind(owner_dest); |
| 672 | } |
| 673 | if (nonzero) { |
| 674 | auto&& vec = m_uninitialized_assignor_grad_oprs; |
| 675 | while (!vec.empty()) { |
| 676 | auto opr = vec.back(); |
| 677 | vec.pop_back(); |
| 678 | opr->init_assignee_info( |
| 679 | m_assignor2info.at(opr->assignor()).assignees, m_grad_virtual_loss); |
| 680 | } |
| 681 | } |
| 682 | return nonzero; |
| 683 | } |
| 684 | |
| 685 | void LoopGrad::GradDesc::init_assignments() { |
| 686 | for (auto&& i : m_fwd_graph_modifier.assignee2info()) { |
no test coverage detected