MCPcopy Create free account
hub / github.com/Oneflow-Inc/oneflow / FwInputArgModifyFn

Function FwInputArgModifyFn

oneflow/user/ops/normalization_op.cpp:75–97  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

73}
74
75Maybe<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
99Maybe<void> FwGetSbpFn(user_op::SbpContext* ctx) {
100 std::vector<user_op::OpArg> split_args;

Callers 1

ModifyInputArgMethod · 0.85

Calls 2

has_inputMethod · 0.45
set_requires_gradMethod · 0.45

Tested by

no test coverage detected