| 74 | : m_opr(opr), m_wrt(wrt), m_grad(grad) {} |
| 75 | |
| 76 | static void make(OperatorNodeBase* opr, VarNode* wrt, VarNode* grad) { |
| 77 | if (ComputingGraphImpl::downcast(wrt->owner_graph()) |
| 78 | ->eager_eval_manager() |
| 79 | .enabled()) |
| 80 | return; |
| 81 | using namespace std::placeholders; |
| 82 | auto checker = std::make_shared<GradShapeChecker>(opr, wrt, grad); |
| 83 | auto func = std::bind(&on_var_shape, checker, _1); |
| 84 | wrt->add_shape_update_callback(grad, func); |
| 85 | grad->add_shape_update_callback(wrt, func); |
| 86 | |
| 87 | if (wrt->shape().ndim && grad->shape().ndim) { |
| 88 | // eager check if shape available |
| 89 | checker->do_on_var_shape(wrt); |
| 90 | checker->do_on_var_shape(grad); |
| 91 | } |
| 92 | } |
| 93 | }; // GradShapeChecker |
| 94 | |
| 95 | struct StaticData { |
nothing calls this directly
no test coverage detected