From 128d5cbd5bb6a9ccd851b34e2f85a1a8f7a2ca23 Mon Sep 17 00:00:00 2001 From: Suryansh Sijwali Date: Fri, 14 Aug 2026 09:08:35 -0400 Subject: [PATCH] Check dim order in the optimized layer_norm as the portable one does --- .../optimized/cpu/op_native_layer_norm.cpp | 10 +++++++ kernels/test/op_native_layer_norm_test.cpp | 26 +++++++++++++++++++ 2 files changed, 36 insertions(+) diff --git a/kernels/optimized/cpu/op_native_layer_norm.cpp b/kernels/optimized/cpu/op_native_layer_norm.cpp index 5fac9faf25e..d47a3e2350a 100644 --- a/kernels/optimized/cpu/op_native_layer_norm.cpp +++ b/kernels/optimized/cpu/op_native_layer_norm.cpp @@ -149,6 +149,16 @@ std::tuple opt_native_layer_norm_out( InvalidArgument, ret_val); + // Only support default dim order for now. + ET_KERNEL_CHECK( + ctx, tensor_is_default_dim_order(input), InvalidArgument, ret_val); + + ET_KERNEL_CHECK( + ctx, + tensors_have_same_dim_order(input, out, mean_out, rstd_out), + InvalidArgument, + ret_val); + Tensor::SizesType mean_rstd_sizes[kTensorDimensionLimit]; size_t mean_rstd_ndim = 0; get_layer_norm_out_target_size( diff --git a/kernels/test/op_native_layer_norm_test.cpp b/kernels/test/op_native_layer_norm_test.cpp index e1345a10354..3429b336580 100644 --- a/kernels/test/op_native_layer_norm_test.cpp +++ b/kernels/test/op_native_layer_norm_test.cpp @@ -452,3 +452,29 @@ TEST_F(OpNativeLayerNormTest, DynamicShapeUnbound) { test_dynamic_shape( {1, 1}, torch::executor::TensorShapeDynamism::DYNAMIC_UNBOUND); } + +TEST_F(OpNativeLayerNormTest, NonDefaultDimOrderDies) { + TensorFactory tf; + + // mean and rstd share the input's rank with the normalized dims set to 1. + // All four are channels-last so the same-dim-order check passes and only the + // default dim order check can reject. + Tensor input = tf.channels_last_like( + tf.make({1, 3, 2, 2}, {1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12})); + Tensor out0 = tf.zeros_channels_last({1, 3, 2, 2}); + Tensor out1 = tf.zeros_channels_last({1, 3, 2, 1}); + Tensor out2 = tf.zeros_channels_last({1, 3, 2, 1}); + const std::vector normalized_shape = {2}; + + ET_EXPECT_KERNEL_FAILURE( + context_, + op_native_layer_norm_out( + input, + IntArrayRef(normalized_shape.data(), normalized_shape.size()), + exec_aten::optional(), + exec_aten::optional(), + 1e-5, + out0, + out1, + out2)); +}