| 73 | } |
| 74 | |
| 75 | Maybe<void> FwInputArgModifyFn(const user_op::GetInputArgModifier& GetInputArgModifierFn, |
| 76 | const user_op::UserOpConfWrapper& conf) { |
| 77 | bool training = true; |
| 78 | if (conf.op_type_name() == "normalization" || conf.op_type_name() == "normalization_add_relu") { |
| 79 | training = conf.attr<bool>("training"); |
| 80 | } |
| 81 | if (conf.has_input("moving_mean", 0)) { |
| 82 | CHECK_OR_RETURN(conf.has_input("moving_variance", 0)); |
| 83 | user_op::InputArgModifier* moving_mean_modifier = GetInputArgModifierFn("moving_mean", 0); |
| 84 | CHECK_OR_RETURN(moving_mean_modifier != nullptr); |
| 85 | moving_mean_modifier->set_is_mutable(training); |
| 86 | moving_mean_modifier->set_requires_grad(false); |
| 87 | user_op::InputArgModifier* moving_variance_modifier = |
| 88 | GetInputArgModifierFn("moving_variance", 0); |
| 89 | CHECK_OR_RETURN(moving_variance_modifier != nullptr); |
| 90 | moving_variance_modifier->set_is_mutable(training); |
| 91 | moving_variance_modifier->set_requires_grad(false); |
| 92 | } else { |
| 93 | CHECK_OR_RETURN(training) |
| 94 | << "Must have moving mean and moving variance for normalization in inference mode."; |
| 95 | } |
| 96 | return Maybe<void>::Ok(); |
| 97 | } |
| 98 | |
| 99 | Maybe<void> FwGetSbpFn(user_op::SbpContext* ctx) { |
| 100 | std::vector<user_op::OpArg> split_args; |
no test coverage detected