From 16eed68a1eea02799a8879a148f1fb17ede59f63 Mon Sep 17 00:00:00 2001 From: Cian Hayes Date: Mon, 8 Feb 2021 17:36:14 -0800 Subject: [PATCH] Fix layer_norm.cc on x86 (#6556) * Fix LayerNromGrad on x86 * PR feedback --- .../orttraining/training_ops/cpu/nn/layer_norm.cc | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/orttraining/orttraining/training_ops/cpu/nn/layer_norm.cc b/orttraining/orttraining/training_ops/cpu/nn/layer_norm.cc index 30e1cfc405..8f131c80d3 100644 --- a/orttraining/orttraining/training_ops/cpu/nn/layer_norm.cc +++ b/orttraining/orttraining/training_ops/cpu/nn/layer_norm.cc @@ -58,8 +58,10 @@ Status LayerNormGrad::Compute(OpKernelContext* op_kernel_context) const Tensor* X = op_kernel_context->Input(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::max()); + ORT_ENFORCE(X_shape.SizeFromDimension(axis) <= std::numeric_limits::max()); + const auto N = static_cast(X_shape.SizeToDimension(axis)); + const auto M = static_cast(X_shape.SizeFromDimension(axis)); ORT_ENFORCE(M != 1); const Tensor* scale = op_kernel_context->Input(input_index++); @@ -146,8 +148,10 @@ Status InvertibleLayerNormGrad::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::max()); + ORT_ENFORCE(X_shape.SizeFromDimension(axis) <= std::numeric_limits::max()); + const auto N = static_cast(X_shape.SizeToDimension(axis)); + const auto M = static_cast(X_shape.SizeFromDimension(axis)); ORT_ENFORCE(M != 1); const auto& scale_shape = scale->Shape();