Fix layer_norm.cc on x86 (#6556)

* Fix LayerNromGrad on x86

* PR feedback
This commit is contained in:
Cian Hayes 2021-02-08 17:36:14 -08:00 committed by GitHub
parent 13d7db9a98
commit 16eed68a1e
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23

View file

@ -58,8 +58,10 @@ Status LayerNormGrad<T, simplified>::Compute(OpKernelContext* op_kernel_context)
const Tensor* X = op_kernel_context->Input<Tensor>(input_index++);
const auto& X_shape = X->Shape();
const auto axis = HandleNegativeAxis(axis_, X_shape.NumDimensions());
const auto N = X_shape.SizeToDimension(axis);
const auto M = X_shape.SizeFromDimension(axis);
ORT_ENFORCE(X_shape.SizeToDimension(axis) <= std::numeric_limits<Eigen::Index>::max());
ORT_ENFORCE(X_shape.SizeFromDimension(axis) <= std::numeric_limits<Eigen::Index>::max());
const auto N = static_cast<Eigen::Index>(X_shape.SizeToDimension(axis));
const auto M = static_cast<Eigen::Index>(X_shape.SizeFromDimension(axis));
ORT_ENFORCE(M != 1);
const Tensor* scale = op_kernel_context->Input<Tensor>(input_index++);
@ -146,8 +148,10 @@ Status InvertibleLayerNormGrad<T>::Compute(OpKernelContext* op_kernel_context) c
const auto& Y_shape = Y_grad->Shape();
const auto& X_shape = Y_shape;
const auto axis = HandleNegativeAxis(axis_, X_shape.NumDimensions());
const auto N = X_shape.SizeToDimension(axis);
const auto M = X_shape.SizeFromDimension(axis);
ORT_ENFORCE(X_shape.SizeToDimension(axis) <= std::numeric_limits<Eigen::Index>::max());
ORT_ENFORCE(X_shape.SizeFromDimension(axis) <= std::numeric_limits<Eigen::Index>::max());
const auto N = static_cast<Eigen::Index>(X_shape.SizeToDimension(axis));
const auto M = static_cast<Eigen::Index>(X_shape.SizeFromDimension(axis));
ORT_ENFORCE(M != 1);
const auto& scale_shape = scale->Shape();