| 803 | |
| 804 | #if MGB_ENABLE_GRAD |
| 805 | MGB_IMPL_OPR_GRAD(CondExecMark) { |
| 806 | if (wrt_idx == opr.input().size() - 1 || !out_grad.at(wrt_idx)) { |
| 807 | return nullptr; |
| 808 | } |
| 809 | using GradMode = CondExecMark::Param::GradMode; |
| 810 | using MergeMode = CondExecMerge::Param::Mode; |
| 811 | MergeMode grad_mode; |
| 812 | SymbolVarArray grad_shapes; |
| 813 | switch (opr.param().grad_mode) { |
| 814 | case GradMode::SUM: |
| 815 | grad_mode = MergeMode::SUM; |
| 816 | grad_shapes.emplace_back(SymbolVar{opr.input(wrt_idx)}.symshape()); |
| 817 | break; |
| 818 | case GradMode::SUM_COND_OUT: |
| 819 | grad_mode = MergeMode::SUM_COND_OUT; |
| 820 | break; |
| 821 | default: |
| 822 | mgb_throw(MegBrainError, "invalid grad_mode"); |
| 823 | } |
| 824 | return CondExecMerge::make_opr( |
| 825 | {out_grad[wrt_idx]}, grad_shapes, {1, grad_mode}, |
| 826 | OperatorNodeConfig{}) |
| 827 | ->output(0); |
| 828 | } |
| 829 | #endif |
| 830 | |
| 831 | /* ============================= CondExecMerge ============================= */ |
nothing calls this directly
no test coverage detected