| 683 | } |
| 684 | |
| 685 | void LoopGrad::GradDesc::init_assignments() { |
| 686 | for (auto&& i : m_fwd_graph_modifier.assignee2info()) { |
| 687 | m_assignor2info[i.second.assignor].assignees.push_back(i.first); |
| 688 | } |
| 689 | |
| 690 | auto grad_trans = [this](VarNode* target, VarNode* wrt, VarNode* grad) { |
| 691 | mgb_assert(target == m_grad_virtual_loss.node()); |
| 692 | auto gnew = AssignorGradOpr::make(grad, wrt); |
| 693 | auto&& d = m_assignor2info.at(wrt); |
| 694 | mgb_assert(!d.grad_opr); |
| 695 | d.grad_opr = &gnew.node()->owner_opr()->cast_final_safe<AssignorGradOpr>(); |
| 696 | m_uninitialized_assignor_grad_oprs.push_back(d.grad_opr); |
| 697 | return gnew.node(); |
| 698 | }; |
| 699 | |
| 700 | for (auto&& i : m_assignor2info) { |
| 701 | cg::add_grad_transformer(i.first, grad_trans); |
| 702 | |
| 703 | for (VarNode* j : i.second.assignees) { |
| 704 | // assignee := assignor, and grads on assignor can be computed if we |
| 705 | // have grads on assignee; assinee is output var, and assignor is |
| 706 | // input var |
| 707 | cg::add_extra_dep_for_grad(i.first, j); |
| 708 | } |
| 709 | } |
| 710 | } |
| 711 | |
| 712 | void LoopGrad::GradDesc::on_sub_graph_func_compile( |
| 713 | ComputingGraph::OutputSpec& out_spec) { |