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

Function BwTensorDescInferFn

oneflow/user/ops/normalization_op.cpp:433–458  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

431namespace {
432
433Maybe<void> BwTensorDescInferFn(user_op::InferContext* ctx) {
434#ifdef WITH_CUDA
435 // assume cudnn is enabled
436 CHECK_GE_OR_RETURN(ctx->Attr<float>("epsilon"), CUDNN_BN_MIN_EPSILON);
437#endif
438 const user_op::TensorDesc& x = ctx->InputTensorDesc("x", 0);
439 const Shape& x_shape = x.shape();
440 const user_op::TensorDesc& dy = ctx->InputTensorDesc("dy", 0);
441 CHECK_EQ_OR_RETURN(dy.shape(), x_shape);
442 if (ctx->has_input("y", 0)) {
443 const user_op::TensorDesc& y = ctx->InputTensorDesc("y", 0);
444 CHECK_EQ_OR_RETURN(y.shape(), x_shape);
445 }
446 *ctx->MutOutputTensorDesc("dx", 0) = x;
447 if (ctx->has_output("addend_diff", 0)) { *ctx->MutOutputTensorDesc("addend_diff", 0) = x; }
448 const Shape param_shape({x_shape.At(ctx->Attr<int32_t>("axis"))});
449 const auto CheckParamTensorDesc = MakeCheckParamTensorDescFn(ctx, param_shape);
450 const auto SetParamTensorDesc = MakeSetParamTensorDescFn(ctx, param_shape);
451 JUST(CheckParamTensorDesc("mean"));
452 JUST(CheckParamTensorDesc("inv_variance"));
453 JUST(CheckParamTensorDesc("gamma"));
454 JUST(CheckParamTensorDesc("beta"));
455 JUST(SetParamTensorDesc("gamma_diff"));
456 JUST(SetParamTensorDesc("beta_diff"));
457 return Maybe<void>::Ok();
458}
459
460Maybe<void> BwDataTypeInferFn(user_op::InferContext* ctx) {
461 const user_op::TensorDesc& x = ctx->InputTensorDesc("x", 0);

Callers 1

Calls 7

MakeSetParamTensorDescFnFunction · 0.85
shapeMethod · 0.45
has_inputMethod · 0.45
MutOutputTensorDescMethod · 0.45
has_outputMethod · 0.45
AtMethod · 0.45

Tested by

no test coverage detected