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();