| 431 | namespace { |
| 432 | |
| 433 | Maybe<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 | |
| 460 | Maybe<void> BwDataTypeInferFn(user_op::InferContext* ctx) { |
| 461 | const user_op::TensorDesc& x = ctx->InputTensorDesc("x", 0); |
no test coverage detected