From 2a43bef7a3c56ce000372da553a59a63b5b49acc Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Thu, 24 Sep 2026 19:41:14 -0700 Subject: [PATCH 01/21] Return 0 from the zero branch of the xlogy/xlog1py/xdivy GPU kernels The MLIR-generated GPU kernels for these ops, and the TF-dialect lowering in lower_tf.td that follows the same pattern, select x itself when x == 0: select(x == 0, x, x / y) That returns -0 for a -0 input. The GPU kernels also flush denormals, so a subnormal x compares equal to zero and comes back unchanged, which is how xdivy(1e-40, 1e-40) returns 1e-40 on GPU (#127913). Select zero instead, as BinaryNoNanPat already does for div_no_nan and as the tf2xla kernels for these ops already do. The complex xlog1py kernel also compared x against 1+0j rather than 0, so xlog1py(1, y) returned 1 for any y, and xlog1py(0, -1) returned nan. Compare against zero. The new tests run only on GPU. On CPU these ops run the Eigen functors in cwise_ops.h, which #127946 covers. --- .../mlir/tensorflow/tests/lower_tf.mlir | 6 +-- .../mlir/tensorflow/transforms/lower_tf.td | 7 +++- .../op_definitions/xdivy.mlir.tmpl | 16 +++---- .../op_definitions/xdivy_cmplx.mlir.tmpl | 16 +++---- .../op_definitions/xlog1py.mlir.tmpl | 16 +++---- .../op_definitions/xlog1py_cmplx.mlir.tmpl | 18 ++++---- .../op_definitions/xlogy.mlir.tmpl | 16 +++---- .../op_definitions/xlogy_cmplx.mlir.tmpl | 16 +++---- tensorflow/python/ops/math_ops_test.py | 42 +++++++++++++++++++ 9 files changed, 99 insertions(+), 54 deletions(-) diff --git a/tensorflow/compiler/mlir/tensorflow/tests/lower_tf.mlir b/tensorflow/compiler/mlir/tensorflow/tests/lower_tf.mlir index a6288e0790cde8..cfa20bfd2055e7 100644 --- a/tensorflow/compiler/mlir/tensorflow/tests/lower_tf.mlir +++ b/tensorflow/compiler/mlir/tensorflow/tests/lower_tf.mlir @@ -1052,7 +1052,7 @@ func.func @xdivy(%lhs: tensor<*xf32>, %rhs: tensor<*xf32>) -> tensor<*xf32> { // CHECK: %[[ZERO:.*]] = "tf.Const"() <{value = dense<0.000000e+00> : tensor}> : () -> tensor // CHECK: %[[IS_ZERO:.*]] = "tf.Equal"(%[[X]], %[[ZERO]]) <{incompatible_shape_error = true}> : (tensor<*xf32>, tensor) -> tensor<*xi1> // CHECK: %[[MUL:.*]] = "tf.Div"(%[[X]], %[[Y]]) : (tensor<*xf32>, tensor<*xf32>) -> tensor<*xf32> - // CHECK: %[[RESULT:.*]] = "tf.SelectV2"(%[[IS_ZERO]], %[[X]], %[[MUL]]) : (tensor<*xi1>, tensor<*xf32>, tensor<*xf32>) -> tensor<*xf32> + // CHECK: %[[RESULT:.*]] = "tf.SelectV2"(%[[IS_ZERO]], %[[ZERO]], %[[MUL]]) : (tensor<*xi1>, tensor, tensor<*xf32>) -> tensor<*xf32> %0 = "tf.Xdivy"(%lhs, %rhs) : (tensor<*xf32>, tensor<*xf32>) -> tensor<*xf32> // CHECK: return %[[RESULT]] func.return %0 : tensor<*xf32> @@ -1065,7 +1065,7 @@ func.func @xlog1py(%lhs: tensor<*xf32>, %rhs: tensor<*xf32>) -> tensor<*xf32> { // CHECK: %[[IS_ZERO:.*]] = "tf.Equal"(%[[X]], %[[ZERO]]) <{incompatible_shape_error = true}> : (tensor<*xf32>, tensor) -> tensor<*xi1> // CHECK: %[[LOG:.*]] = "tf.Log1p"(%[[Y]]) : (tensor<*xf32>) -> tensor<*xf32> // CHECK: %[[MUL:.*]] = "tf.Mul"(%[[X]], %[[LOG]]) : (tensor<*xf32>, tensor<*xf32>) -> tensor<*xf32> - // CHECK: %[[RESULT:.*]] = "tf.SelectV2"(%[[IS_ZERO]], %[[X]], %[[MUL]]) : (tensor<*xi1>, tensor<*xf32>, tensor<*xf32>) -> tensor<*xf32> + // CHECK: %[[RESULT:.*]] = "tf.SelectV2"(%[[IS_ZERO]], %[[ZERO]], %[[MUL]]) : (tensor<*xi1>, tensor, tensor<*xf32>) -> tensor<*xf32> %0 = "tf.Xlog1py"(%lhs, %rhs) : (tensor<*xf32>, tensor<*xf32>) -> tensor<*xf32> // CHECK: return %[[RESULT]] func.return %0 : tensor<*xf32> @@ -1078,7 +1078,7 @@ func.func @xlogy(%lhs: tensor<*xf32>, %rhs: tensor<*xf32>) -> tensor<*xf32> { // CHECK: %[[IS_ZERO:.*]] = "tf.Equal"(%[[X]], %[[ZERO]]) <{incompatible_shape_error = true}> : (tensor<*xf32>, tensor) -> tensor<*xi1> // CHECK: %[[LOG:.*]] = "tf.Log"(%[[Y]]) : (tensor<*xf32>) -> tensor<*xf32> // CHECK: %[[MUL:.*]] = "tf.Mul"(%[[X]], %[[LOG]]) : (tensor<*xf32>, tensor<*xf32>) -> tensor<*xf32> - // CHECK: %[[RESULT:.*]] = "tf.SelectV2"(%[[IS_ZERO]], %[[X]], %[[MUL]]) : (tensor<*xi1>, tensor<*xf32>, tensor<*xf32>) -> tensor<*xf32> + // CHECK: %[[RESULT:.*]] = "tf.SelectV2"(%[[IS_ZERO]], %[[ZERO]], %[[MUL]]) : (tensor<*xi1>, tensor, tensor<*xf32>) -> tensor<*xf32> %0 = "tf.Xlogy"(%lhs, %rhs) : (tensor<*xf32>, tensor<*xf32>) -> tensor<*xf32> // CHECK: return %[[RESULT]] func.return %0 : tensor<*xf32> diff --git a/tensorflow/compiler/mlir/tensorflow/transforms/lower_tf.td b/tensorflow/compiler/mlir/tensorflow/transforms/lower_tf.td index 1061d564f51afc..caa458c859f373 100644 --- a/tensorflow/compiler/mlir/tensorflow/transforms/lower_tf.td +++ b/tensorflow/compiler/mlir/tensorflow/transforms/lower_tf.td @@ -479,17 +479,20 @@ def LowerScatterNdOp : // Xdivy, Xlog1p and Xlogy op patterns. //===----------------------------------------------------------------------===// +// Selects zero rather than $x when $x == 0, like BinaryNoNanPat above: $x can +// compare equal to zero without being +0, as -0 does and as a subnormal does +// when denormals are flushed. class BinaryXopyPat : Pat< From, (TF_SelectV2Op (TF_EqualOp $x, - (TF_ConstOp + (TF_ConstOp:$zero (GetScalarOfType<0> $x) ), /*incompatible_shape_error*/ConstBoolAttrTrue ), - $x, + $zero, To )>; diff --git a/tensorflow/core/kernels/mlir_generated/op_definitions/xdivy.mlir.tmpl b/tensorflow/core/kernels/mlir_generated/op_definitions/xdivy.mlir.tmpl index 429b017f986c02..4ad024b1489be5 100644 --- a/tensorflow/core/kernels/mlir_generated/op_definitions/xdivy.mlir.tmpl +++ b/tensorflow/core/kernels/mlir_generated/op_definitions/xdivy.mlir.tmpl @@ -22,7 +22,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %20 = tensor.reshape %arg1(%from_elements) : (tensor<*xelem_type>, tensor<1xindex>) -> tensor %21 = chlo.broadcast_compare %19, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %22 = chlo.broadcast_divide %19, %20 : (tensor, tensor) -> tensor - %23 = chlo.broadcast_select %21, %19, %22 : (tensor, tensor, tensor) -> tensor + %23 = chlo.broadcast_select %21, %5, %22 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %23 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -35,7 +35,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %23 = tensor.reshape %arg1(%c_empty) : (tensor<*xelem_type>, tensor<0xindex>) -> tensor %24 = chlo.broadcast_compare %22, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %25 = chlo.broadcast_divide %22, %23 : (tensor, tensor) -> tensor - %26 = chlo.broadcast_select %24, %22, %25 : (tensor, tensor, tensor) -> tensor + %26 = chlo.broadcast_select %24, %5, %25 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %26 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -48,7 +48,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %26 = tensor.reshape %arg1(%from_elements) : (tensor<*xelem_type>, tensor<1xindex>) -> tensor %27 = chlo.broadcast_compare %25, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %28 = chlo.broadcast_divide %25, %26 : (tensor, tensor) -> tensor - %29 = chlo.broadcast_select %27, %25, %28 : (tensor, tensor, tensor) -> tensor + %29 = chlo.broadcast_select %27, %5, %28 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %29 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -67,7 +67,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %33 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<1xindex>) -> tensor %34 = chlo.broadcast_compare %31, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %35 = chlo.broadcast_divide %31, %33 : (tensor, tensor) -> tensor - %36 = chlo.broadcast_select %34, %31, %35 : (tensor, tensor, tensor) -> tensor + %36 = chlo.broadcast_select %34, %5, %35 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %36 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -81,7 +81,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %35 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<2xindex>) -> tensor %36 = chlo.broadcast_compare %33, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %37 = chlo.broadcast_divide %33, %35 : (tensor, tensor) -> tensor - %38 = chlo.broadcast_select %36, %33, %37 : (tensor, tensor, tensor) -> tensor + %38 = chlo.broadcast_select %36, %5, %37 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %38 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -95,7 +95,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %37 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<3xindex>) -> tensor %38 = chlo.broadcast_compare %35, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %39 = chlo.broadcast_divide %35, %37 : (tensor, tensor) -> tensor - %40 = chlo.broadcast_select %38, %35, %39 : (tensor, tensor, tensor) -> tensor + %40 = chlo.broadcast_select %38, %5, %39 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %40 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -109,7 +109,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %39 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<4xindex>) -> tensor %40 = chlo.broadcast_compare %37, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %41 = chlo.broadcast_divide %37, %39 : (tensor, tensor) -> tensor - %42 = chlo.broadcast_select %40, %37, %41 : (tensor, tensor, tensor) -> tensor + %42 = chlo.broadcast_select %40, %5, %41 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %42 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -123,7 +123,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %40 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<5xindex>) -> tensor %41 = chlo.broadcast_compare %38, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %42 = chlo.broadcast_divide %38, %40 : (tensor, tensor) -> tensor - %43 = chlo.broadcast_select %41, %38, %42 : (tensor, tensor, tensor) -> tensor + %43 = chlo.broadcast_select %41, %5, %42 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %43 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } diff --git a/tensorflow/core/kernels/mlir_generated/op_definitions/xdivy_cmplx.mlir.tmpl b/tensorflow/core/kernels/mlir_generated/op_definitions/xdivy_cmplx.mlir.tmpl index ae2ce94f00e092..f0c3ef523f4b03 100644 --- a/tensorflow/core/kernels/mlir_generated/op_definitions/xdivy_cmplx.mlir.tmpl +++ b/tensorflow/core/kernels/mlir_generated/op_definitions/xdivy_cmplx.mlir.tmpl @@ -22,7 +22,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %20 = tensor.reshape %arg1(%from_elements) : (tensor<*xelem_type>, tensor<1xindex>) -> tensor %21 = chlo.broadcast_compare %19, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %22 = chlo.broadcast_divide %19, %20 : (tensor, tensor) -> tensor - %23 = chlo.broadcast_select %21, %19, %22 : (tensor, tensor, tensor) -> tensor + %23 = chlo.broadcast_select %21, %5, %22 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %23 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -35,7 +35,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %23 = tensor.reshape %arg1(%c_empty) : (tensor<*xelem_type>, tensor<0xindex>) -> tensor %24 = chlo.broadcast_compare %22, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %25 = chlo.broadcast_divide %22, %23 : (tensor, tensor) -> tensor - %26 = chlo.broadcast_select %24, %22, %25 : (tensor, tensor, tensor) -> tensor + %26 = chlo.broadcast_select %24, %5, %25 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %26 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -48,7 +48,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %26 = tensor.reshape %arg1(%from_elements) : (tensor<*xelem_type>, tensor<1xindex>) -> tensor %27 = chlo.broadcast_compare %25, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %28 = chlo.broadcast_divide %25, %26 : (tensor, tensor) -> tensor - %29 = chlo.broadcast_select %27, %25, %28 : (tensor, tensor, tensor) -> tensor + %29 = chlo.broadcast_select %27, %5, %28 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %29 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -67,7 +67,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %33 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<1xindex>) -> tensor %34 = chlo.broadcast_compare %31, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %35 = chlo.broadcast_divide %31, %33 : (tensor, tensor) -> tensor - %36 = chlo.broadcast_select %34, %31, %35 : (tensor, tensor, tensor) -> tensor + %36 = chlo.broadcast_select %34, %5, %35 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %36 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -81,7 +81,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %35 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<2xindex>) -> tensor %36 = chlo.broadcast_compare %33, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %37 = chlo.broadcast_divide %33, %35 : (tensor, tensor) -> tensor - %38 = chlo.broadcast_select %36, %33, %37 : (tensor, tensor, tensor) -> tensor + %38 = chlo.broadcast_select %36, %5, %37 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %38 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -95,7 +95,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %37 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<3xindex>) -> tensor %38 = chlo.broadcast_compare %35, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %39 = chlo.broadcast_divide %35, %37 : (tensor, tensor) -> tensor - %40 = chlo.broadcast_select %38, %35, %39 : (tensor, tensor, tensor) -> tensor + %40 = chlo.broadcast_select %38, %5, %39 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %40 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -109,7 +109,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %39 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<4xindex>) -> tensor %40 = chlo.broadcast_compare %37, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %41 = chlo.broadcast_divide %37, %39 : (tensor, tensor) -> tensor - %42 = chlo.broadcast_select %40, %37, %41 : (tensor, tensor, tensor) -> tensor + %42 = chlo.broadcast_select %40, %5, %41 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %42 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -123,7 +123,7 @@ func.func @Xdivy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %40 = tensor.reshape %arg1(%cast_0) : (tensor<*xelem_type>, tensor<5xindex>) -> tensor %41 = chlo.broadcast_compare %38, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %42 = chlo.broadcast_divide %38, %40 : (tensor, tensor) -> tensor - %43 = chlo.broadcast_select %41, %38, %42 : (tensor, tensor, tensor) -> tensor + %43 = chlo.broadcast_select %41, %5, %42 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %43 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } diff --git a/tensorflow/core/kernels/mlir_generated/op_definitions/xlog1py.mlir.tmpl b/tensorflow/core/kernels/mlir_generated/op_definitions/xlog1py.mlir.tmpl index 97d76e45a604ef..3fdee68fa12651 100644 --- a/tensorflow/core/kernels/mlir_generated/op_definitions/xlog1py.mlir.tmpl +++ b/tensorflow/core/kernels/mlir_generated/op_definitions/xlog1py.mlir.tmpl @@ -23,7 +23,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %21 = chlo.broadcast_compare %19, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %22 = mhlo.log_plus_one %20 : tensor %23 = chlo.broadcast_multiply %19, %22 : (tensor, tensor) -> tensor - %24 = chlo.broadcast_select %21, %19, %23 : (tensor, tensor, tensor) -> tensor + %24 = chlo.broadcast_select %21, %5, %23 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %24 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -37,7 +37,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %24 = chlo.broadcast_compare %22, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %25 = mhlo.log_plus_one %23 : tensor %26 = chlo.broadcast_multiply %22, %25 : (tensor, tensor) -> tensor - %27 = chlo.broadcast_select %24, %22, %26 : (tensor, tensor, tensor) -> tensor + %27 = chlo.broadcast_select %24, %5, %26 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %27 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -51,7 +51,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %27 = chlo.broadcast_compare %25, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %28 = mhlo.log_plus_one %26 : tensor %29 = chlo.broadcast_multiply %25, %28 : (tensor, tensor) -> tensor - %30 = chlo.broadcast_select %27, %25, %29 : (tensor, tensor, tensor) -> tensor + %30 = chlo.broadcast_select %27, %5, %29 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %30 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -71,7 +71,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %34 = chlo.broadcast_compare %31, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %35 = mhlo.log_plus_one %33 : tensor %36 = chlo.broadcast_multiply %31, %35 : (tensor, tensor) -> tensor - %37 = chlo.broadcast_select %34, %31, %36 : (tensor, tensor, tensor) -> tensor + %37 = chlo.broadcast_select %34, %5, %36 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %37 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -86,7 +86,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %36 = chlo.broadcast_compare %33, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %37 = mhlo.log_plus_one %35 : tensor %38 = chlo.broadcast_multiply %33, %37 : (tensor, tensor) -> tensor - %39 = chlo.broadcast_select %36, %33, %38 : (tensor, tensor, tensor) -> tensor + %39 = chlo.broadcast_select %36, %5, %38 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %39 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -101,7 +101,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %38 = chlo.broadcast_compare %35, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %39 = mhlo.log_plus_one %37 : tensor %40 = chlo.broadcast_multiply %35, %39 : (tensor, tensor) -> tensor - %41 = chlo.broadcast_select %38, %35, %40 : (tensor, tensor, tensor) -> tensor + %41 = chlo.broadcast_select %38, %5, %40 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %41 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -116,7 +116,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %40 = chlo.broadcast_compare %37, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %41 = mhlo.log_plus_one %39 : tensor %42 = chlo.broadcast_multiply %37, %41 : (tensor, tensor) -> tensor - %43 = chlo.broadcast_select %40, %37, %42 : (tensor, tensor, tensor) -> tensor + %43 = chlo.broadcast_select %40, %5, %42 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %43 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -131,7 +131,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %41 = chlo.broadcast_compare %38, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %42 = mhlo.log_plus_one %40 : tensor %43 = chlo.broadcast_multiply %38, %42 : (tensor, tensor) -> tensor - %44 = chlo.broadcast_select %41, %38, %43 : (tensor, tensor, tensor) -> tensor + %44 = chlo.broadcast_select %41, %5, %43 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %44 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } diff --git a/tensorflow/core/kernels/mlir_generated/op_definitions/xlog1py_cmplx.mlir.tmpl b/tensorflow/core/kernels/mlir_generated/op_definitions/xlog1py_cmplx.mlir.tmpl index 8c884d20688e15..73a03a295f6c18 100644 --- a/tensorflow/core/kernels/mlir_generated/op_definitions/xlog1py_cmplx.mlir.tmpl +++ b/tensorflow/core/kernels/mlir_generated/op_definitions/xlog1py_cmplx.mlir.tmpl @@ -9,7 +9,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %c2 = arith.constant 2 : index %4 = shape.const_shape [1] : tensor<1xindex> %c1 = arith.constant 1 : index - %5 = mhlo.constant dense<(1.000000e+00,0.000000e+00)> : tensor + %5 = mhlo.constant dense<(0.000000e+00,0.000000e+00)> : tensor %6 = shape.shape_of %arg0 : tensor<*xelem_type> -> tensor %7 = shape.shape_of %arg1 : tensor<*xelem_type> -> tensor %8 = shape.num_elements %6 : tensor -> index @@ -23,7 +23,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %21 = chlo.broadcast_compare %19, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %22 = mhlo.log_plus_one %20 : tensor %23 = chlo.broadcast_multiply %19, %22 : (tensor, tensor) -> tensor - %24 = chlo.broadcast_select %21, %19, %23 : (tensor, tensor, tensor) -> tensor + %24 = chlo.broadcast_select %21, %5, %23 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %24 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -37,7 +37,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %24 = chlo.broadcast_compare %22, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %25 = mhlo.log_plus_one %23 : tensor %26 = chlo.broadcast_multiply %22, %25 : (tensor, tensor) -> tensor - %27 = chlo.broadcast_select %24, %22, %26 : (tensor, tensor, tensor) -> tensor + %27 = chlo.broadcast_select %24, %5, %26 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %27 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -51,7 +51,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %27 = chlo.broadcast_compare %25, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %28 = mhlo.log_plus_one %26 : tensor %29 = chlo.broadcast_multiply %25, %28 : (tensor, tensor) -> tensor - %30 = chlo.broadcast_select %27, %25, %29 : (tensor, tensor, tensor) -> tensor + %30 = chlo.broadcast_select %27, %5, %29 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %30 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -71,7 +71,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %34 = chlo.broadcast_compare %31, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %35 = mhlo.log_plus_one %33 : tensor %36 = chlo.broadcast_multiply %31, %35 : (tensor, tensor) -> tensor - %37 = chlo.broadcast_select %34, %31, %36 : (tensor, tensor, tensor) -> tensor + %37 = chlo.broadcast_select %34, %5, %36 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %37 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -86,7 +86,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %36 = chlo.broadcast_compare %33, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %37 = mhlo.log_plus_one %35 : tensor %38 = chlo.broadcast_multiply %33, %37 : (tensor, tensor) -> tensor - %39 = chlo.broadcast_select %36, %33, %38 : (tensor, tensor, tensor) -> tensor + %39 = chlo.broadcast_select %36, %5, %38 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %39 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -101,7 +101,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %38 = chlo.broadcast_compare %35, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %39 = mhlo.log_plus_one %37 : tensor %40 = chlo.broadcast_multiply %35, %39 : (tensor, tensor) -> tensor - %41 = chlo.broadcast_select %38, %35, %40 : (tensor, tensor, tensor) -> tensor + %41 = chlo.broadcast_select %38, %5, %40 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %41 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -116,7 +116,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %40 = chlo.broadcast_compare %37, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %41 = mhlo.log_plus_one %39 : tensor %42 = chlo.broadcast_multiply %37, %41 : (tensor, tensor) -> tensor - %43 = chlo.broadcast_select %40, %37, %42 : (tensor, tensor, tensor) -> tensor + %43 = chlo.broadcast_select %40, %5, %42 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %43 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -131,7 +131,7 @@ func.func @Xlog1py_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %a %41 = chlo.broadcast_compare %38, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %42 = mhlo.log_plus_one %40 : tensor %43 = chlo.broadcast_multiply %38, %42 : (tensor, tensor) -> tensor - %44 = chlo.broadcast_select %41, %38, %43 : (tensor, tensor, tensor) -> tensor + %44 = chlo.broadcast_select %41, %5, %43 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %44 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } diff --git a/tensorflow/core/kernels/mlir_generated/op_definitions/xlogy.mlir.tmpl b/tensorflow/core/kernels/mlir_generated/op_definitions/xlogy.mlir.tmpl index 06c20a8226e00d..bb3429db57322b 100644 --- a/tensorflow/core/kernels/mlir_generated/op_definitions/xlogy.mlir.tmpl +++ b/tensorflow/core/kernels/mlir_generated/op_definitions/xlogy.mlir.tmpl @@ -23,7 +23,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %21 = chlo.broadcast_compare %19, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %22 = mhlo.log %20 : tensor %23 = chlo.broadcast_multiply %19, %22 : (tensor, tensor) -> tensor - %24 = chlo.broadcast_select %21, %19, %23 : (tensor, tensor, tensor) -> tensor + %24 = chlo.broadcast_select %21, %5, %23 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %24 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -37,7 +37,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %24 = chlo.broadcast_compare %22, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %25 = mhlo.log %23 : tensor %26 = chlo.broadcast_multiply %22, %25 : (tensor, tensor) -> tensor - %27 = chlo.broadcast_select %24, %22, %26 : (tensor, tensor, tensor) -> tensor + %27 = chlo.broadcast_select %24, %5, %26 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %27 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -51,7 +51,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %27 = chlo.broadcast_compare %25, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %28 = mhlo.log %26 : tensor %29 = chlo.broadcast_multiply %25, %28 : (tensor, tensor) -> tensor - %30 = chlo.broadcast_select %27, %25, %29 : (tensor, tensor, tensor) -> tensor + %30 = chlo.broadcast_select %27, %5, %29 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %30 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -71,7 +71,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %34 = chlo.broadcast_compare %31, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %35 = mhlo.log %33 : tensor %36 = chlo.broadcast_multiply %31, %35 : (tensor, tensor) -> tensor - %37 = chlo.broadcast_select %34, %31, %36 : (tensor, tensor, tensor) -> tensor + %37 = chlo.broadcast_select %34, %5, %36 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %37 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -86,7 +86,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %36 = chlo.broadcast_compare %33, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %37 = mhlo.log %35 : tensor %38 = chlo.broadcast_multiply %33, %37 : (tensor, tensor) -> tensor - %39 = chlo.broadcast_select %36, %33, %38 : (tensor, tensor, tensor) -> tensor + %39 = chlo.broadcast_select %36, %5, %38 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %39 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -101,7 +101,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %38 = chlo.broadcast_compare %35, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %39 = mhlo.log %37 : tensor %40 = chlo.broadcast_multiply %35, %39 : (tensor, tensor) -> tensor - %41 = chlo.broadcast_select %38, %35, %40 : (tensor, tensor, tensor) -> tensor + %41 = chlo.broadcast_select %38, %5, %40 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %41 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -116,7 +116,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %40 = chlo.broadcast_compare %37, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %41 = mhlo.log %39 : tensor %42 = chlo.broadcast_multiply %37, %41 : (tensor, tensor) -> tensor - %43 = chlo.broadcast_select %40, %37, %42 : (tensor, tensor, tensor) -> tensor + %43 = chlo.broadcast_select %40, %5, %42 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %43 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -131,7 +131,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %41 = chlo.broadcast_compare %38, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %42 = mhlo.log %40 : tensor %43 = chlo.broadcast_multiply %38, %42 : (tensor, tensor) -> tensor - %44 = chlo.broadcast_select %41, %38, %43 : (tensor, tensor, tensor) -> tensor + %44 = chlo.broadcast_select %41, %5, %43 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %44 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } diff --git a/tensorflow/core/kernels/mlir_generated/op_definitions/xlogy_cmplx.mlir.tmpl b/tensorflow/core/kernels/mlir_generated/op_definitions/xlogy_cmplx.mlir.tmpl index 172ba62d10c28c..b61f6bc2367ba9 100644 --- a/tensorflow/core/kernels/mlir_generated/op_definitions/xlogy_cmplx.mlir.tmpl +++ b/tensorflow/core/kernels/mlir_generated/op_definitions/xlogy_cmplx.mlir.tmpl @@ -23,7 +23,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %21 = chlo.broadcast_compare %19, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %22 = mhlo.log %20 : tensor %23 = chlo.broadcast_multiply %19, %22 : (tensor, tensor) -> tensor - %24 = chlo.broadcast_select %21, %19, %23 : (tensor, tensor, tensor) -> tensor + %24 = chlo.broadcast_select %21, %5, %23 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %24 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -37,7 +37,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %24 = chlo.broadcast_compare %22, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %25 = mhlo.log %23 : tensor %26 = chlo.broadcast_multiply %22, %25 : (tensor, tensor) -> tensor - %27 = chlo.broadcast_select %24, %22, %26 : (tensor, tensor, tensor) -> tensor + %27 = chlo.broadcast_select %24, %5, %26 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %27 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -51,7 +51,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %27 = chlo.broadcast_compare %25, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %28 = mhlo.log %26 : tensor %29 = chlo.broadcast_multiply %25, %28 : (tensor, tensor) -> tensor - %30 = chlo.broadcast_select %27, %25, %29 : (tensor, tensor, tensor) -> tensor + %30 = chlo.broadcast_select %27, %5, %29 : (tensor, tensor, tensor) -> tensor %cast = tensor.cast %30 : tensor to tensor<*xelem_type> scf.yield %cast : tensor<*xelem_type> } else { @@ -71,7 +71,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %34 = chlo.broadcast_compare %31, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %35 = mhlo.log %33 : tensor %36 = chlo.broadcast_multiply %31, %35 : (tensor, tensor) -> tensor - %37 = chlo.broadcast_select %34, %31, %36 : (tensor, tensor, tensor) -> tensor + %37 = chlo.broadcast_select %34, %5, %36 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %37 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -86,7 +86,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %36 = chlo.broadcast_compare %33, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %37 = mhlo.log %35 : tensor %38 = chlo.broadcast_multiply %33, %37 : (tensor, tensor) -> tensor - %39 = chlo.broadcast_select %36, %33, %38 : (tensor, tensor, tensor) -> tensor + %39 = chlo.broadcast_select %36, %5, %38 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %39 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -101,7 +101,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %38 = chlo.broadcast_compare %35, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %39 = mhlo.log %37 : tensor %40 = chlo.broadcast_multiply %35, %39 : (tensor, tensor) -> tensor - %41 = chlo.broadcast_select %38, %35, %40 : (tensor, tensor, tensor) -> tensor + %41 = chlo.broadcast_select %38, %5, %40 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %41 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -116,7 +116,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %40 = chlo.broadcast_compare %37, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %41 = mhlo.log %39 : tensor %42 = chlo.broadcast_multiply %37, %41 : (tensor, tensor) -> tensor - %43 = chlo.broadcast_select %40, %37, %42 : (tensor, tensor, tensor) -> tensor + %43 = chlo.broadcast_select %40, %5, %42 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %43 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } else { @@ -131,7 +131,7 @@ func.func @Xlogy_platform_elem_type_output_type(%arg0: tensor<*xelem_type>, %arg %41 = chlo.broadcast_compare %38, %5 {comparison_direction = #chlo} : (tensor, tensor) -> tensor %42 = mhlo.log %40 : tensor %43 = chlo.broadcast_multiply %38, %42 : (tensor, tensor) -> tensor - %44 = chlo.broadcast_select %41, %38, %43 : (tensor, tensor, tensor) -> tensor + %44 = chlo.broadcast_select %41, %5, %43 : (tensor, tensor, tensor) -> tensor %cast_1 = tensor.cast %44 : tensor to tensor<*xelem_type> scf.yield %cast_1 : tensor<*xelem_type> } diff --git a/tensorflow/python/ops/math_ops_test.py b/tensorflow/python/ops/math_ops_test.py index 84b81cb1151e0f..cfb5f339db3398 100644 --- a/tensorflow/python/ops/math_ops_test.py +++ b/tensorflow/python/ops/math_ops_test.py @@ -1104,6 +1104,48 @@ def testBasic(self): self.assertAllEqual(zeros * x, tf_result_reverseargs) +@test_util.run_all_in_graph_and_eager_modes +class XopsGpuKernelTest(test_util.TensorFlowTestCase): + """The zero branch of the MLIR-generated GPU kernels for xlogy and friends. + + On CPU these ops run the Eigen functors in cwise_ops.h instead. + """ + + @test_util.run_gpu_only + def testNegativeZeroGivesPositiveZero(self): + for op, y in ((math_ops.xlogy, 2.), (math_ops.xlog1py, 1.), + (math_ops.xdivy, 3.)): + for dtype in [dtypes.float16, dtypes.float32, dtypes.float64]: + x = constant_op.constant(np.full((16,), -0.), dtype=dtype) + y_t = constant_op.constant(np.full((16,), y), dtype=dtype) + with test_util.force_gpu(): + result = self.evaluate(op(x, y_t)) + self.assertAllEqual(result, np.zeros(16)) + self.assertFalse(np.signbit(result).any()) + + @test_util.run_gpu_only + def testSubnormalIsNotReturned(self): + # The GPU kernels flush denormals, so a subnormal x compares equal to zero + # and the result is 0; without flushing, x / x is 1. It must never be x. + for dtype in [dtypes.float16, dtypes.float32, dtypes.float64]: + tiny = np.finfo(dtype.as_numpy_dtype).tiny / 4 + x = constant_op.constant(np.full((16,), tiny), dtype=dtype) + with test_util.force_gpu(): + result = self.evaluate(math_ops.xdivy(x, x)) + self.assertTrue(np.isin(result, [0., 1.]).all(), result) + + @test_util.run_gpu_only + def testComplexXlog1pyComparesXAgainstZero(self): + # The complex kernel compared x against 1 rather than 0, so xlog1py(1, y) + # returned 1 and xlog1py(0, -1) returned nan. + for dtype in [dtypes.complex64, dtypes.complex128]: + x = constant_op.constant([1., 0.], dtype=dtype) + y = constant_op.constant([1., -1.], dtype=dtype) + with test_util.force_gpu(): + result = self.evaluate(math_ops.xlog1py(x, y)) + self.assertAllClose(result, [np.log(2.), 0.]) + + @test_util.run_all_in_graph_and_eager_modes class XlogyTest(test_util.TensorFlowTestCase): From d85401017bbf62b3494f0c2e56ccd1fde70baa15 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Fri, 25 Sep 2026 09:56:28 -0700 Subject: [PATCH 02/21] Update the xlogy/xlog1py/xdivy GPU test baselines for +0 The C++ baselines in gpu_binary_ops_test.cc returned x when x == 0, which is what the kernels did before this change. For x = -0 that is -0, which ExpectStrictlyEqual tells apart from the +0 the kernels now return, so the Half tests failed. Return zero, as the kernels do. --- .../core/kernels/mlir_generated/gpu_binary_ops_test.cc | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/tensorflow/core/kernels/mlir_generated/gpu_binary_ops_test.cc b/tensorflow/core/kernels/mlir_generated/gpu_binary_ops_test.cc index 51087da47053c3..a4ff86d4389e3a 100644 --- a/tensorflow/core/kernels/mlir_generated/gpu_binary_ops_test.cc +++ b/tensorflow/core/kernels/mlir_generated/gpu_binary_ops_test.cc @@ -1496,7 +1496,7 @@ TEST_F(BinaryOpsTest, SubUint32SpecialCases) { template T baseline_xlogy(T x, T y) { - return x == T(0) ? x : x * std::log(y); + return x == T(0) ? T(0) : x * std::log(y); } GENERATE_DEFAULT_TESTS_2(Xlogy, /*test_name=*/Half, Eigen::half, float, @@ -1519,12 +1519,13 @@ GENERATE_DEFAULT_TESTS(Xlogy, /*test_name=*/Complex128, std::complex, template T baseline_xlog1py(T x, T y) { - return x == T(0) ? x : x * std::log1p(y); + return x == T(0) ? T(0) : x * std::log1p(y); } template std::complex baseline_xlog1py(std::complex x, std::complex y) { - return x == std::complex(0) ? x : x * std::log(std::complex(1) + y); + return x == std::complex(0) ? std::complex(0) + : x * std::log(std::complex(1) + y); } GENERATE_DEFAULT_TESTS_2(Xlog1py, /*test_name=*/Half, Eigen::half, float, @@ -1569,7 +1570,7 @@ GENERATE_DEFAULT_TESTS_WITH_SPECIFIC_INPUT_VALUES( template T baseline_xdivy(T x, T y) { - return x == T(0) ? x : x / y; + return x == T(0) ? T(0) : x / y; } GENERATE_DEFAULT_TESTS_2(Xdivy, /*test_name=*/Half, Eigen::half, float, From f424878592142052d0f921852b4fa4228e92bbe1 Mon Sep 17 00:00:00 2001 From: Siqiao Wu Date: Tue, 29 Sep 2026 15:57:00 -0700 Subject: [PATCH 03/21] Set intra op threadpool in ynn param when using NanoRt Executable. PiperOrigin-RevId: 990583743 --- third_party/xla/xla/backends/cpu/nanort/BUILD | 5 + .../backends/cpu/nanort/nanort_client_test.cc | 104 ++++++++++++++++++ .../backends/cpu/nanort/nanort_executable.cc | 23 +++- 3 files changed, 129 insertions(+), 3 deletions(-) diff --git a/third_party/xla/xla/backends/cpu/nanort/BUILD b/third_party/xla/xla/backends/cpu/nanort/BUILD index 2cca231ba060f9..0448b37e73de2b 100644 --- a/third_party/xla/xla/backends/cpu/nanort/BUILD +++ b/third_party/xla/xla/backends/cpu/nanort/BUILD @@ -69,6 +69,7 @@ xla_cc_test( "//xla:literal_util", "//xla:shape_util", "//xla:xla_data_proto_cc", + "//xla:xla_proto_cc", "//xla/backends/cpu:alignment", "//xla/ffi", "//xla/ffi:execution_context", @@ -85,7 +86,9 @@ xla_cc_test( "//xla/pjrt/plugin/xla_cpu:xla_cpu_pjrt_client", "//xla/runtime:device_id", "//xla/service:computation_placer", + "//xla/service:hlo_module_config", "//xla/service/cpu:cpu_aot_compilation_result", + "//xla/service/cpu:cpu_executable", "//xla/tsl/concurrency:async_value", "//xla/tsl/lib/core:status_test_util", "//xla/tsl/platform:logging", @@ -121,6 +124,8 @@ cc_library( "//xla/backends/cpu/runtime:function_library", "//xla/backends/cpu/runtime:thread_pool_task_runner", "//xla/backends/cpu/runtime:thunk", + "//xla/backends/cpu/runtime/ynnpack:ynn_interop", + "//xla/backends/cpu/runtime/ynnpack:ynn_threadpool", "//xla/ffi:execution_context", "//xla/hlo/ir:hlo", "//xla/runtime:device_id", diff --git a/third_party/xla/xla/backends/cpu/nanort/nanort_client_test.cc b/third_party/xla/xla/backends/cpu/nanort/nanort_client_test.cc index ec23293b9a8a8b..3369fd167cc244 100644 --- a/third_party/xla/xla/backends/cpu/nanort/nanort_client_test.cc +++ b/third_party/xla/xla/backends/cpu/nanort/nanort_client_test.cc @@ -17,8 +17,10 @@ limitations under the License. #include +#include #include #include +#include #include #include #include @@ -52,6 +54,8 @@ limitations under the License. #include "xla/runtime/device_id.h" #include "xla/service/computation_placer.h" #include "xla/service/cpu/cpu_aot_compilation_result.h" +#include "xla/service/cpu/cpu_executable.h" +#include "xla/service/hlo_module_config.h" #include "xla/shape_util.h" #include "xla/tsl/concurrency/async_value_ref.h" #include "xla/tsl/lib/core/status_test_util.h" @@ -59,6 +63,7 @@ limitations under the License. #include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "xla/tsl/platform/test_benchmark.h" +#include "xla/xla.pb.h" #include "xla/xla_data.pb.h" #include "tsl/platform/casts.h" @@ -307,6 +312,105 @@ ENTRY test_module { EXPECT_EQ(result_span[0], expected_result); } +// Eigen thread pool that counts the number of scheduled tasks. +class CountingThreadPool : public Eigen::ThreadPoolInterface { + public: + explicit CountingThreadPool(int num_threads) : pool_(num_threads) {} + + void Schedule(std::function fn) override { + num_scheduled_.fetch_add(1, std::memory_order_relaxed); + pool_.Schedule(std::move(fn)); + } + + void ScheduleWithHint(std::function fn, int start, + int limit) override { + num_scheduled_.fetch_add(1, std::memory_order_relaxed); + pool_.ScheduleWithHint(std::move(fn), start, limit); + } + + int NumThreads() const override { return pool_.NumThreads(); } + int CurrentThreadId() const override { return pool_.CurrentThreadId(); } + + int64_t num_scheduled() const { + return num_scheduled_.load(std::memory_order_relaxed); + } + + private: + Eigen::ThreadPool pool_; + std::atomic num_scheduled_{0}; +}; + +// Regression test: YNN fusions executed via NanoRtExecutable must use the +// intra-op thread pool instead of running single-threaded in the caller thread. +TEST_P(NanoRtClientTest, YnnFusionUsesIntraOpThreadPool) { + constexpr absl::string_view hlo = R"( + HloModule ynn_dot + + ENTRY e { + p0 = f32[256,512] parameter(0) + p1 = f32[512,512] parameter(1) + ROOT dot = f32[256,512] dot(p0, p1), + lhs_contracting_dims={1}, rhs_contracting_dims={0} + } + )"; + + TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnUnverifiedModule(hlo)); + XlaComputation computation(module->ToProto()); + + NanoRtClient client([](HloModuleConfig& config) { + DebugOptions debug_options = config.debug_options(); + debug_options.clear_xla_cpu_experimental_ynn_fusion_type(); + debug_options.add_xla_cpu_experimental_ynn_fusion_type( + DebugOptions::LIBRARY_FUSION_TYPE_INDIVIDUAL_DOT); + config.set_debug_options(debug_options); + }); + TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr executable, + client.Compile(computation)); + + if (GetParam()) { + TF_ASSERT_OK_AND_ASSIGN(auto exported, client.Export(executable.get())); + auto* aot_compilation_result = + absl::down_cast(exported.get()); + TF_ASSERT_OK_AND_ASSIGN( + executable, NanoRtExecutable::Create(aot_compilation_result->proto(), + executable->program_shape())); + } + + // Make sure the dot was actually offloaded to YNNPACK, otherwise this test + // doesn't test anything. + auto* cpu_executable = + absl::down_cast(executable->executable()); + ASSERT_TRUE(cpu_executable->has_ynn_fusions()); + + Array2D lhs(256, 512, 1.0f); + Array2D rhs(512, 512, 1.0f); + Array2D result(256, 512, 0.0f); + + Arguments arguments = { + {lhs.data(), static_cast(lhs.num_elements())}, + {rhs.data(), static_cast(rhs.num_elements())}}; + Results results = { + {result.data(), static_cast(result.num_elements())}}; + NanoRtExecutable::ManagedTemp<32> temp(executable->temp_buffer_size()); + + CountingThreadPool tp(4); + Eigen::ThreadPoolDevice device(&tp, tp.NumThreads()); + + NanoRtExecutable::ExecuteOptions execute_options; + execute_options.set_intra_op_thread_pool(&device); + auto event = executable->Execute(arguments, results, temp, execute_options); + tsl::BlockUntilReady(event); + + ASSERT_TRUE(event.IsConcrete()); + EXPECT_EQ(result(0, 0), 512.0f); + EXPECT_EQ(result(255, 511), 512.0f); + + // A single YNN fusion thunk is executed inline in the caller thread by the + // thunk executor, so all tasks scheduled on the intra-op thread pool come + // from the YNNPACK runtime parallelizing the dot. + EXPECT_GT(tp.num_scheduled(), 0); +} + TEST_P(NanoRtClientTest, CompileAndRunPartitionAndReplicaIdInstructions) { constexpr absl::string_view hlo = R"( HloModule replica-and-partition-id diff --git a/third_party/xla/xla/backends/cpu/nanort/nanort_executable.cc b/third_party/xla/xla/backends/cpu/nanort/nanort_executable.cc index e1cf35472a7dd5..01591e370cb8e5 100644 --- a/third_party/xla/xla/backends/cpu/nanort/nanort_executable.cc +++ b/third_party/xla/xla/backends/cpu/nanort/nanort_executable.cc @@ -34,6 +34,8 @@ limitations under the License. #include "xla/backends/cpu/runtime/function_library.h" #include "xla/backends/cpu/runtime/thread_pool_task_runner.h" #include "xla/backends/cpu/runtime/thunk.h" +#include "xla/backends/cpu/runtime/ynnpack/ynn_interop.h" +#include "xla/backends/cpu/runtime/ynnpack/ynn_threadpool.h" #include "xla/executable_run_options.h" #include "xla/ffi/execution_context.h" #include "xla/hlo/ir/hlo_module.h" @@ -404,7 +406,8 @@ tsl::AsyncValueRef NanoRtExecutable::Execute( struct ExecutionContext { ExecutionContext(cpu::BufferAllocations::Buffers buffers, FunctionLibrary* function_library, - const ExecuteOptions& options) + const ExecuteOptions& options, + std::optional ynn_params) : allocations(std::move(buffers)), execute_params(Thunk::ExecuteParams{function_library, &allocations, /*xfeed=*/nullptr, @@ -417,15 +420,19 @@ tsl::AsyncValueRef NanoRtExecutable::Execute( options.device_assignment(), /*collectives=*/nullptr), custom_call_execute_params( RunId(options.launch_id()), options.local_device_id().value(), - options.intra_op_thread_pool(), options.ffi_context()) { + options.intra_op_thread_pool(), options.ffi_context()), + ynn_params(std::move(ynn_params)) { execute_params.collective_params = &collective_execute_params; execute_params.custom_call_params = &custom_call_execute_params; + execute_params.ynn_params = + this->ynn_params.has_value() ? &*this->ynn_params : nullptr; } cpu::BufferAllocations allocations; Thunk::ExecuteParams execute_params; Thunk::CollectiveExecuteParams collective_execute_params; Thunk::CustomCallExecuteParams custom_call_execute_params; + std::optional ynn_params; }; // Do a heap allocation if we're running with a thread pool, using @@ -434,8 +441,18 @@ tsl::AsyncValueRef NanoRtExecutable::Execute( // allocation when it is not required. if (options.intra_op_thread_pool() || options.ffi_context() || options.device_assignment()) { + // Prepare for executing YNNPACK fusions. Without a YNN threadpool, YNNPACK + // fusions run single-threaded in the caller thread. + std::optional ynn_params; + if (executable->has_ynn_fusions() && options.intra_op_thread_pool()) { + ABSL_ASSIGN_OR_RETURN(YnnThreadpool ynn_threadpool, + CreateYnnThreadpool(options.intra_op_thread_pool())); + ynn_params.emplace(std::move(ynn_threadpool)); + } + auto execution_context = std::make_unique( - std::move(buffers), executable->function_library(), options); + std::move(buffers), executable->function_library(), options, + std::move(ynn_params)); auto execute_event = executable->thunks().Execute(execution_context->execute_params); From 059acf0462fe1a178a0936afdab3c5a3537beed1 Mon Sep 17 00:00:00 2001 From: Antonio Sanchez Date: Tue, 29 Sep 2026 16:10:07 -0700 Subject: [PATCH 04/21] Automated Code Change PiperOrigin-RevId: 990590739 --- tensorflow/core/framework/tensor_testutil_test.cc | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/tensorflow/core/framework/tensor_testutil_test.cc b/tensorflow/core/framework/tensor_testutil_test.cc index 7baf57b2d433dd..6eb0288a609560 100644 --- a/tensorflow/core/framework/tensor_testutil_test.cc +++ b/tensorflow/core/framework/tensor_testutil_test.cc @@ -226,23 +226,23 @@ TEST(TensorTestUtilTest, ExpectTensorCloseHalf) { EXPECT_TRUE(IsClose(static_cast(1.0f), static_cast(1.0f), 0.0, 0.0)); EXPECT_FALSE(IsClose(static_cast(1.0f), static_cast(1.1f), 0.0, 0.0)); - // Epsilon: 0 00010 0000000000 -> 2^-13 = 0.0001220703125 - // Default Tolerance: 0 00100 0100000000 -> 5/2^13 = 0.0006103515625 + // Epsilon: 0 00101 0000000000 -> 2^-10 = 0.0009765625 + // Default Tolerance: 5 * 2^-10 = 0.0048828125 // 1.234 -> 0 01111 0011110000 -> 1264/2^10 = 1.234375 // 1.233 -> 0 01111 0011101111 -> 1263/2^10 = 1.2333984375 // 1.235 -> 0 01111 0011110001 -> 1265/2^10 = 1.2353515625 // 1.232 -> 0 01111 0011101110 -> 1262/2^10 = 1.232421875 // 1.236 -> 0 01111 0011110010 -> 1266/2^10 = 1.236328125 - // 1/2^10 = 0.0009765625E - // Threshold = 0.0013637542724609375 + // 1.200 -> 1229/2^10 = 1.2001953125 + // 1.260 -> 1290/2^10 = 1.259765625 EXPECT_TRUE(IsClose(static_cast(1.234f), static_cast(1.234f))); EXPECT_TRUE(IsClose(static_cast(1.234f), static_cast(1.233f))); EXPECT_TRUE(IsClose(static_cast(1.234f), static_cast(1.235f))); - // Diff = 0.001953125 - EXPECT_FALSE(IsClose(static_cast(1.234f), static_cast(1.232f))); - EXPECT_FALSE(IsClose(static_cast(1.234f), static_cast(1.236f))); + // Diff exceeds default tolerance (atol + rtol * abs(x) ~ 0.0109) + EXPECT_FALSE(IsClose(static_cast(1.234f), static_cast(1.200f))); + EXPECT_FALSE(IsClose(static_cast(1.234f), static_cast(1.260f))); EXPECT_TRUE( IsClose(static_cast(1.234f), static_cast(1.232f), 8e-4f, 1e-3f)); EXPECT_TRUE( From 0a3b7c4a9a28fe43e4ecce9bb82af5c4967faf67 Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 16:27:06 -0700 Subject: [PATCH 05/21] Protect XNNPACK runtime creation and destruction with workspace_mutex_. `Subgraph::Prepare` and `Subgraph::Invoke` acquire `Delegate::workspace_mutex_` to serialize access to the shared `xnn_workspace` and its intrusive `first_user` linked list of `xnn_runtime` instances. However, `Subgraph::Create` (`xnn_create_runtime_v4`), `Subgraph::~Subgraph` (`xnn_delete_runtime`), and `Delegate::~Delegate` (`xnn_release_workspace`) previously mutated `workspace->first_user` and `workspace->ref_count` or freed the `xnn_runtime` without holding `workspace_mutex_`. Share `workspace_mutex_` via `std::shared_ptr` between `Delegate` and `Subgraph` (matching the ref-counted lifetime of `xnn_workspace`), acquire `workspace_mutex_` around `xnn_create_runtime_v4` in `Subgraph::Create`, before resetting `runtime_` in `Subgraph::~Subgraph`, and before resetting `workspace_` in `Delegate::~Delegate`, clean up `runtime_ptr` under `workspace_mutex_` if `StopBuildStep()` fails, and log an error and return `kTfLiteError` if `runtime_` is null in `Subgraph::Prepare` and `Subgraph::Invoke`. PiperOrigin-RevId: 990599327 --- tensorflow/lite/delegates/xnnpack/BUILD | 6 + .../lite/delegates/xnnpack/delegate_test.cc | 134 ++++++++++++++++++ .../delegates/xnnpack/xnnpack_delegate.cc | 65 ++++++--- 3 files changed, 183 insertions(+), 22 deletions(-) diff --git a/tensorflow/lite/delegates/xnnpack/BUILD b/tensorflow/lite/delegates/xnnpack/BUILD index 93dcc92d204914..d039493d440b3b 100644 --- a/tensorflow/lite/delegates/xnnpack/BUILD +++ b/tensorflow/lite/delegates/xnnpack/BUILD @@ -1520,8 +1520,14 @@ cc_test( deps = [ ":test_main", ":xnnpack_delegate_test_mode", + "//tensorflow/lite:framework", + "//tensorflow/lite:schema_fbs_version", "//tensorflow/lite/c:c_api_types", + "//tensorflow/lite/core:framework", + "//tensorflow/lite/core/kernels:builtin_ops", + "//tensorflow/lite/schema:schema_fbs", "@com_google_googletest//:gtest", + "@flatbuffers//:runtime_cc", "@pthreadpool", ], ) diff --git a/tensorflow/lite/delegates/xnnpack/delegate_test.cc b/tensorflow/lite/delegates/xnnpack/delegate_test.cc index fc31ca077c2d3c..033448693bc08c 100644 --- a/tensorflow/lite/delegates/xnnpack/delegate_test.cc +++ b/tensorflow/lite/delegates/xnnpack/delegate_test.cc @@ -13,15 +13,94 @@ See the License for the specific language governing permissions and limitations under the License. ==============================================================================*/ +#include +#include +#include #include +#include // NOLINT(build/c++11) +#include #include +#include "flatbuffers/buffer.h" // from @flatbuffers +#include "flatbuffers/flatbuffer_builder.h" // from @flatbuffers #include "pthreadpool.h" // from @pthreadpool #include "tensorflow/lite/c/c_api_types.h" +#include "tensorflow/lite/core/interpreter_builder.h" +#include "tensorflow/lite/core/kernels/register.h" #include "tensorflow/lite/delegates/xnnpack/xnnpack_delegate.h" +#include "tensorflow/lite/interpreter.h" +#include "tensorflow/lite/schema/schema_generated.h" +#include "tensorflow/lite/version.h" namespace tflite { namespace xnnpack { +namespace { + +std::vector CreateMultiOpModel() { + flatbuffers::FlatBufferBuilder builder; + const std::array, 2> operator_codes{{ + CreateOperatorCode(builder, BuiltinOperator_ABS), + CreateOperatorCode(builder, BuiltinOperator_NEG), + }}; + + const std::array, 1> buffers{{ + CreateBuffer(builder, builder.CreateVector({})), + }}; + + const std::array shape{{1, 16}}; + const std::array, 3> tensors{{ + CreateTensor(builder, + builder.CreateVector(shape.data(), shape.size()), + TensorType_FLOAT32), + CreateTensor(builder, + builder.CreateVector(shape.data(), shape.size()), + TensorType_FLOAT32), + CreateTensor(builder, + builder.CreateVector(shape.data(), shape.size()), + TensorType_FLOAT32), + }}; + + const std::array op0_inputs{{0}}; + const std::array op0_outputs{{1}}; + const std::array op1_inputs{{1}}; + const std::array op1_outputs{{2}}; + const std::array, 2> ops{{ + CreateOperator( + builder, /*opcode_index=*/0, + builder.CreateVector(op0_inputs.data(), op0_inputs.size()), + builder.CreateVector(op0_outputs.data(), + op0_outputs.size())), + CreateOperator( + builder, /*opcode_index=*/1, + builder.CreateVector(op1_inputs.data(), op1_inputs.size()), + builder.CreateVector(op1_outputs.data(), + op1_outputs.size())), + }}; + + const std::array subgraph_inputs{{0}}; + const std::array subgraph_outputs{{2}}; + flatbuffers::Offset subgraph = CreateSubGraph( + builder, builder.CreateVector(tensors.data(), tensors.size()), + builder.CreateVector(subgraph_inputs.data(), + subgraph_inputs.size()), + builder.CreateVector(subgraph_outputs.data(), + subgraph_outputs.size()), + builder.CreateVector(ops.data(), ops.size())); + + flatbuffers::Offset model_buffer = CreateModel( + builder, TFLITE_SCHEMA_VERSION, + builder.CreateVector(operator_codes.data(), operator_codes.size()), + builder.CreateVector(&subgraph, 1), + builder.CreateString("Multi-op model"), + builder.CreateVector(buffers.data(), buffers.size())); + + builder.Finish(model_buffer); + + return std::vector(builder.GetBufferPointer(), + builder.GetBufferPointer() + builder.GetSize()); +} + +} // namespace TEST(Delegate, CreateWithoutParams) { std::unique_ptr @@ -60,5 +139,60 @@ TEST(Delegate, GetThreadPool) { ASSERT_EQ(2, pthreadpool_get_threads_count(threadpool)); } +TEST(Delegate, ConcurrentInvokeAndSubgraphCreateDestroy) { + std::vector buffer = CreateMultiOpModel(); + const Model* model = GetModel(buffer.data()); + ::tflite::ops::builtin::BuiltinOpResolverWithoutDefaultDelegates resolver; + + std::unique_ptr + xnnpack_delegate(TfLiteXNNPackDelegateCreate(nullptr), + TfLiteXNNPackDelegateDelete); + + std::unique_ptr invoke_interpreter; + ASSERT_EQ(InterpreterBuilder(model, resolver)(&invoke_interpreter), + kTfLiteOk); + ASSERT_NE(invoke_interpreter, nullptr); + ASSERT_EQ(invoke_interpreter->AllocateTensors(), kTfLiteOk); + ASSERT_EQ(invoke_interpreter->ModifyGraphWithDelegate(xnnpack_delegate.get()), + kTfLiteOk); + ASSERT_EQ(invoke_interpreter->Invoke(), kTfLiteOk); + + std::atomic stop{false}; + std::thread invoke_thread([&]() { + int step = 1; + while (!stop.load(std::memory_order_relaxed)) { + EXPECT_EQ(invoke_interpreter->ResizeInputTensor( + invoke_interpreter->inputs()[0], {1, 16 + step}), + kTfLiteOk); + EXPECT_EQ(invoke_interpreter->AllocateTensors(), kTfLiteOk); + EXPECT_EQ(invoke_interpreter->Invoke(), kTfLiteOk); + ++step; + } + }); + + std::thread create_destroy_thread([&]() { + for (int i = 0; i < 50; ++i) { + std::unique_ptr interpreter; + EXPECT_EQ(InterpreterBuilder(model, resolver)(&interpreter), kTfLiteOk); + ASSERT_NE(interpreter, nullptr); + EXPECT_EQ(interpreter->ResizeInputTensor(interpreter->inputs()[0], + {1, 16 * (i + 1)}), + kTfLiteOk); + EXPECT_EQ(interpreter->AllocateTensors(), kTfLiteOk); + EXPECT_EQ(interpreter->ModifyGraphWithDelegate(xnnpack_delegate.get()), + kTfLiteOk); + EXPECT_EQ(interpreter->Invoke(), kTfLiteOk); + } + stop.store(true, std::memory_order_relaxed); + }); + + create_destroy_thread.join(); + invoke_thread.join(); + + std::thread destroy_thread([&]() { invoke_interpreter.reset(); }); + xnnpack_delegate.reset(); + destroy_thread.join(); +} + } // namespace xnnpack } // namespace tflite diff --git a/tensorflow/lite/delegates/xnnpack/xnnpack_delegate.cc b/tensorflow/lite/delegates/xnnpack/xnnpack_delegate.cc index 9a5773b985ae43..aca6f01e1c97d5 100644 --- a/tensorflow/lite/delegates/xnnpack/xnnpack_delegate.cc +++ b/tensorflow/lite/delegates/xnnpack/xnnpack_delegate.cc @@ -740,6 +740,13 @@ class Delegate { } } + ~Delegate() { + if (workspace_mutex_ != nullptr && workspace_ != nullptr) { + std::lock_guard lock(*workspace_mutex_); + workspace_.reset(); + } + } + TfLiteIntArray* PrepareOpsToDelegate(TfLiteContext* context, TfLiteIntArray** moe_ops_to_delegate); TfLiteDelegate* tflite_delegate() { return &delegate_; } @@ -935,7 +942,7 @@ class Delegate { nullptr, &xnn_release_workspace}; TfLiteXNNPackDelegateOptions options_{}; - std::mutex workspace_mutex_; + std::shared_ptr workspace_mutex_ = std::make_shared(); // If no weight cache is provided and a cache is set in the delegate options, // this will be used as a weight cache. @@ -1492,12 +1499,19 @@ class Subgraph { return nullptr; } } - status = xnn_create_runtime_v4(subgraph.get(), delegate.weights_cache(), - delegate.workspace(), delegate.threadpool(), - flags, &runtime_ptr); + { + std::lock_guard lock(*delegate.workspace_mutex_); + status = xnn_create_runtime_v4( + subgraph.get(), delegate.weights_cache(), delegate.workspace(), + delegate.threadpool(), flags, &runtime_ptr); + } if (delegate.weight_cache_provider_->IsActive() && delegate.weight_cache_provider_->CanStartBuildStep()) { if (!delegate.weight_cache_provider_->StopBuildStep()) { + if (runtime_ptr != nullptr) { + std::lock_guard lock(*delegate.workspace_mutex_); + xnn_delete_runtime(runtime_ptr); + } TF_LITE_KERNEL_LOG(context, "XNNPack delegate failed to stop cache build step."); return nullptr; @@ -1514,12 +1528,16 @@ class Subgraph { } TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node, - bool enable_subgraph_reshaping, Delegate* delegate) { + bool enable_subgraph_reshaping) { if (moe_kernel_ != nullptr) { return moe_kernel_->Prepare(context); } - std::lock_guard lock(delegate->workspace_mutex_); + std::lock_guard lock(*workspace_mutex_); + if (runtime_ == nullptr) { + TF_LITE_KERNEL_LOG(context, "XNNPACK runtime is null."); + return kTfLiteError; + } tflite::Subgraph* this_subgraph = reinterpret_cast(context->impl_); @@ -1594,13 +1612,16 @@ class Subgraph { return kTfLiteOk; } - TfLiteStatus Invoke(TfLiteContext* context, bool enable_subgraph_reshaping, - Delegate* delegate) { + TfLiteStatus Invoke(TfLiteContext* context, bool enable_subgraph_reshaping) { if (moe_kernel_ != nullptr) { return moe_kernel_->Invoke(context); } - std::lock_guard lock(delegate->workspace_mutex_); + std::lock_guard lock(*workspace_mutex_); + if (runtime_ == nullptr) { + TF_LITE_KERNEL_LOG(context, "XNNPACK runtime is null."); + return kTfLiteError; + } tflite::Subgraph* this_subgraph = reinterpret_cast(context->impl_); @@ -7154,7 +7175,12 @@ class Subgraph { return enable_subgraph_reshaping_; } - inline Delegate* GetDelegate() const { return delegate_; } + ~Subgraph() { + if (workspace_mutex_ != nullptr && runtime_ != nullptr) { + std::lock_guard lock(*workspace_mutex_); + runtime_.reset(); + } + } private: Subgraph(Delegate& delegate, xnn_runtime_t runtime, @@ -7168,7 +7194,7 @@ class Subgraph { tflite_tensor_to_xnnpack_(std::move(tflite_tensor_to_xnnpack)), resources_(delegate.local_id_to_resources_), enable_subgraph_reshaping_(delegate.enable_subgraph_reshaping()), - delegate_(&delegate) { + workspace_mutex_(delegate.workspace_mutex_) { for (int t : externals) { externals_[t] = nullptr; } @@ -7177,10 +7203,9 @@ class Subgraph { Subgraph(Delegate& delegate, std::unique_ptr moe_kernel) : runtime_(nullptr, &xnn_delete_runtime), - moe_kernel_(std::move(moe_kernel)) { - enable_subgraph_reshaping_ = delegate.enable_subgraph_reshaping(); - delegate_ = &delegate; - } + moe_kernel_(std::move(moe_kernel)), + enable_subgraph_reshaping_(delegate.enable_subgraph_reshaping()), + workspace_mutex_(delegate.workspace_mutex_) {} // Keep track of expanded scales for shared tensors to manage their lifetime. // Must be declared before runtime_ so it outlives runtime_ during @@ -7212,7 +7237,7 @@ class Subgraph { // data pointer to nullptr, and XNNPACK requires valid data pointers. char dummy_data_{0}; bool enable_subgraph_reshaping_ = false; - Delegate* delegate_; + std::shared_ptr workspace_mutex_; }; TfLiteIntArray* Delegate::PrepareOpsToDelegate( @@ -7732,9 +7757,7 @@ TfLiteStatus SubgraphPrepare(TfLiteContext* context, TfLiteNode* node) { } Subgraph* subgraph = static_cast(node->user_data); - return static_cast(node->user_data) - ->Prepare(context, node, subgraph->EnableSubgraphReshaping(), - subgraph->GetDelegate()); + return subgraph->Prepare(context, node, subgraph->EnableSubgraphReshaping()); } TfLiteStatus SubgraphInvoke(TfLiteContext* context, TfLiteNode* node) { @@ -7743,9 +7766,7 @@ TfLiteStatus SubgraphInvoke(TfLiteContext* context, TfLiteNode* node) { } Subgraph* subgraph = static_cast(node->user_data); - return static_cast(node->user_data) - ->Invoke(context, subgraph->EnableSubgraphReshaping(), - subgraph->GetDelegate()); + return subgraph->Invoke(context, subgraph->EnableSubgraphReshaping()); } void SubgraphFree(TfLiteContext* context, void* buffer) { From 4be465e255a6ab92ed436fde4ee2c6c8ace53198 Mon Sep 17 00:00:00 2001 From: Emilio Cota Date: Tue, 29 Sep 2026 16:38:33 -0700 Subject: [PATCH 06/21] [xla:cpu] ir_compiler: remove obsolete comment about AllowFPOpFusion Integrate cl/983398077 (2e332453d92) removed the callers mentioned in the comment. Remove it. PiperOrigin-RevId: 990604835 --- third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc | 6 ------ 1 file changed, 6 deletions(-) diff --git a/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc index b99fbcce5cbe15..05e845e75ef9f4 100644 --- a/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc +++ b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc @@ -472,12 +472,6 @@ llvm::Error IrCompiler::RunIrPasses(llvm::Module& module, // Must run after all optimization passes: middle-end passes behave // differently on instructions that already carry `contract`. - // - // TODO(b/560320144): `AllowFPOpFusion = Fast` is deliberately still set in - // service/cpu/cpu_aot_loader.cc:53, tools/hlo_opt/cpu_opt.cc:217, - // backends/cpu/testlib/kernel_runner.cc:132 and - // service/cpu/ir_emitter_test.cc:258. Drop those once the upstream change - // has landed. llvm_ir::SetAllowContractOnFpArithmetic(module); // Sanitizer instrumentation must be the last IR transformation. From d06416566f81e2945bc3032c36e3da776aa3b4a7 Mon Sep 17 00:00:00 2001 From: Toli Yevtushenko Date: Tue, 29 Sep 2026 16:44:18 -0700 Subject: [PATCH 07/21] Modernize manual status checks, deprecated TF_* macros, and StrCat error builders PiperOrigin-RevId: 990607199 --- third_party/xla/xla/BUILD | 1 + third_party/xla/xla/hlo/transforms/BUILD | 1 + .../xla/xla/hlo/transforms/collectives/BUILD | 1 + .../all_reduce_reassociate_test.cc | 10 ++-- .../collective_permute_decomposer.cc | 7 +-- .../reduce_scatter_combiner_test.cc | 12 ++--- .../reduce_scatter_reassociate_test.cc | 8 ++-- ...emory_placement_to_internal_annotations.cc | 4 +- .../xla/hlo/transforms/hlo_module_stitcher.cc | 27 ++++++----- .../hlo/transforms/host_offload_legalize.cc | 46 +++++++++---------- .../xla/xla/hlo/transforms/host_offloader.cc | 25 +++++----- .../all_gather_pad_ds_simplifier_test.cc | 8 ++-- .../all_gather_permuted_ds_simplifier_test.cc | 10 ++-- .../simplifiers/reduce_window_rewriter.cc | 5 +- .../simplifiers/unflatten_call_graph.cc | 13 +++--- third_party/xla/xla/util.h | 21 ++++++--- third_party/xla/xla/util_test.cc | 20 ++++++++ 17 files changed, 118 insertions(+), 101 deletions(-) diff --git a/third_party/xla/xla/BUILD b/third_party/xla/xla/BUILD index 3da1bb89bd9bc2..7cdc2385ccba07 100644 --- a/third_party/xla/xla/BUILD +++ b/third_party/xla/xla/BUILD @@ -392,6 +392,7 @@ xla_cc_test( "@com_google_absl//absl/algorithm:container", "@com_google_absl//absl/container:inlined_vector", "@com_google_absl//absl/log:check", + "@com_google_absl//absl/status", "@com_google_absl//absl/strings", "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/types:span", diff --git a/third_party/xla/xla/hlo/transforms/BUILD b/third_party/xla/xla/hlo/transforms/BUILD index 1be14d085b13c6..998bafb4644adc 100644 --- a/third_party/xla/xla/hlo/transforms/BUILD +++ b/third_party/xla/xla/hlo/transforms/BUILD @@ -726,6 +726,7 @@ cc_library( hdrs = ["hlo_module_stitcher.h"], deps = [ "//xla:shape_util", + "//xla:util", "//xla/hlo/ir:hlo", "//xla/hlo/pass:hlo_pass", "@com_google_absl//absl/cleanup", diff --git a/third_party/xla/xla/hlo/transforms/collectives/BUILD b/third_party/xla/xla/hlo/transforms/collectives/BUILD index 51a74a2a0043dd..2ddc19e7f1ef4b 100644 --- a/third_party/xla/xla/hlo/transforms/collectives/BUILD +++ b/third_party/xla/xla/hlo/transforms/collectives/BUILD @@ -786,6 +786,7 @@ cc_library( deps = [ ":collective_permute_cycle", "//xla:shape_util", + "//xla:util", "//xla:xla_data_proto_cc", "//xla:xla_proto_cc", "//xla/hlo/ir:hlo", diff --git a/third_party/xla/xla/hlo/transforms/collectives/all_reduce_reassociate_test.cc b/third_party/xla/xla/hlo/transforms/collectives/all_reduce_reassociate_test.cc index db928a30aefae5..275a4f13876303 100644 --- a/third_party/xla/xla/hlo/transforms/collectives/all_reduce_reassociate_test.cc +++ b/third_party/xla/xla/hlo/transforms/collectives/all_reduce_reassociate_test.cc @@ -49,12 +49,10 @@ class AllReduceSimplifierTest : public HloHardwareIndependentTestBase { absl::string_view hlo_module, bool expect_change, bool reassociate_converted_ar = false) { ABSL_ASSIGN_OR_RETURN(auto module, ParseAndReturnVerifiedModule(hlo_module)); - auto changed = - AllReduceReassociate(reassociate_converted_ar).Run(module.get()); - if (!changed.ok()) { - return changed.status(); - } - EXPECT_EQ(changed.value(), expect_change); + ABSL_ASSIGN_OR_RETURN( + bool changed, + AllReduceReassociate(reassociate_converted_ar).Run(module.get())); + EXPECT_EQ(changed, expect_change); return absl::StatusOr>(std::move(module)); } diff --git a/third_party/xla/xla/hlo/transforms/collectives/collective_permute_decomposer.cc b/third_party/xla/xla/hlo/transforms/collectives/collective_permute_decomposer.cc index 83991e68a13546..1e0360bff14e92 100644 --- a/third_party/xla/xla/hlo/transforms/collectives/collective_permute_decomposer.cc +++ b/third_party/xla/xla/hlo/transforms/collectives/collective_permute_decomposer.cc @@ -45,6 +45,7 @@ limitations under the License. #include "xla/service/source_target_pairs.h" #include "xla/shape.h" #include "xla/shape_util.h" +#include "xla/util.h" #include "xla/xla.pb.h" #include "xla/xla_data.pb.h" @@ -191,9 +192,9 @@ static absl::StatusOr DecomposeCollectivePermute( ABSL_RETURN_IF_ERROR(recv_done->AddControlDependencyTo(send)); break; default: - return absl::InvalidArgumentError( - absl::StrCat("Unsupported pipeline parallelism opt level: ", - pipeline_parallelism_opt_level)); + return InvalidArgumentStrCat( + "Unsupported pipeline parallelism opt level: ", + pipeline_parallelism_opt_level); } if (!pipeline_decision.empty()) { diff --git a/third_party/xla/xla/hlo/transforms/collectives/reduce_scatter_combiner_test.cc b/third_party/xla/xla/hlo/transforms/collectives/reduce_scatter_combiner_test.cc index 1f5c7ef3666628..ab72d787f434c8 100644 --- a/third_party/xla/xla/hlo/transforms/collectives/reduce_scatter_combiner_test.cc +++ b/third_party/xla/xla/hlo/transforms/collectives/reduce_scatter_combiner_test.cc @@ -54,17 +54,15 @@ class ReduceScatterCombinerTest : public HloHardwareIndependentTestBase { VLOG(1) << "Before running ReduceScatterCombiner: " << ReduceScatterCount(module.get()) << " reduce-scatter ops"; - auto changed = ReduceScatterCombiner(byte_threshold, count_threshold, - combine_by_dim, combine_while_loops) - .Run(module.get()); - if (!changed.ok()) { - return changed.status(); - } + ABSL_ASSIGN_OR_RETURN(bool changed, + ReduceScatterCombiner(byte_threshold, count_threshold, + combine_by_dim, combine_while_loops) + .Run(module.get())); VLOG(1) << "After running ReduceScatterCombiner: " << ReduceScatterCount(module.get()) << " reduce-scatter ops"; - EXPECT_EQ(changed.value(), expect_change); + EXPECT_EQ(changed, expect_change); return absl::StatusOr>(std::move(module)); } diff --git a/third_party/xla/xla/hlo/transforms/collectives/reduce_scatter_reassociate_test.cc b/third_party/xla/xla/hlo/transforms/collectives/reduce_scatter_reassociate_test.cc index 5c16cf6bcce57b..01fad55557c273 100644 --- a/third_party/xla/xla/hlo/transforms/collectives/reduce_scatter_reassociate_test.cc +++ b/third_party/xla/xla/hlo/transforms/collectives/reduce_scatter_reassociate_test.cc @@ -42,11 +42,9 @@ class ReduceScatterReassociateTest : public HloHardwareIndependentTestBase { absl::StatusOr> RunPass( absl::string_view hlo_module, bool expect_change) { ABSL_ASSIGN_OR_RETURN(auto module, ParseAndReturnVerifiedModule(hlo_module)); - auto changed = ReduceScatterReassociate().Run(module.get()); - if (!changed.ok()) { - return changed.status(); - } - EXPECT_EQ(changed.value(), expect_change); + ABSL_ASSIGN_OR_RETURN(bool changed, + ReduceScatterReassociate().Run(module.get())); + EXPECT_EQ(changed, expect_change); return absl::StatusOr>(std::move(module)); } diff --git a/third_party/xla/xla/hlo/transforms/convert_memory_placement_to_internal_annotations.cc b/third_party/xla/xla/hlo/transforms/convert_memory_placement_to_internal_annotations.cc index f84fcaa1d443ab..bf21d2f0afed34 100644 --- a/third_party/xla/xla/hlo/transforms/convert_memory_placement_to_internal_annotations.cc +++ b/third_party/xla/xla/hlo/transforms/convert_memory_placement_to_internal_annotations.cc @@ -48,8 +48,8 @@ absl::StatusOr GetCustomCallTarget( if (external_annotation == memory_annotations::kMemoryTargetPinnedDevice) { return memory_annotations::kPinToDeviceCustomCallTarget; } - return absl::InvalidArgumentError( - absl::StrCat("Invalid external annotation: ", external_annotation)); + return InvalidArgumentStrCat("Invalid external annotation: ", + external_annotation); } absl::StatusOr diff --git a/third_party/xla/xla/hlo/transforms/hlo_module_stitcher.cc b/third_party/xla/xla/hlo/transforms/hlo_module_stitcher.cc index af0a0c80dddce4..a0cc112979b516 100644 --- a/third_party/xla/xla/hlo/transforms/hlo_module_stitcher.cc +++ b/third_party/xla/xla/hlo/transforms/hlo_module_stitcher.cc @@ -33,6 +33,7 @@ limitations under the License. #include "xla/hlo/ir/hlo_opcode.h" #include "xla/shape.h" #include "xla/shape_util.h" +#include "xla/util.h" namespace xla { @@ -45,9 +46,9 @@ absl::StatusOr HloModuleStitcher::RunImpl( } if (visiting_modules_.contains(module)) { - return absl::InternalError( - absl::StrCat("Circular dependency detected in submodule stitching: ", - module->name())); + return InternalStrCat( + "Circular dependency detected in submodule stitching: ", + module->name()); } if (visited_modules_.contains(module)) { @@ -72,8 +73,7 @@ absl::StatusOr HloModuleStitcher::RunImpl( std::string sub_module_name = inst->raw_backend_config_string(); auto it = optimized_modules_.find(sub_module_name); if (it == optimized_modules_.end()) { - return absl::NotFoundError( - absl::StrCat("Sub-module ", sub_module_name, " not found")); + return NotFoundStrCat("Sub-module ", sub_module_name, " not found"); } HloModule* sub_module = it->second; @@ -85,10 +85,9 @@ absl::StatusOr HloModuleStitcher::RunImpl( HloComputation* sub_entry = sub_module->entry_computation(); if (inst->operand_count() != sub_entry->num_parameters()) { - return absl::InvalidArgumentError(absl::StrCat( + return InvalidArgumentStrCat( "Operand count mismatch: custom call has ", inst->operand_count(), - " operands but sub-module expects ", - sub_entry->num_parameters())); + " operands but sub-module expects ", sub_entry->num_parameters()); } HloCloneContext context(module); @@ -105,10 +104,10 @@ absl::StatusOr HloModuleStitcher::RunImpl( cloned_sub_entry->parameter_instruction(i)->shape(); if (!ShapeUtil::Equal(operand->shape(), expected_shape)) { if (!ShapeUtil::Compatible(operand->shape(), expected_shape)) { - return absl::InvalidArgumentError(absl::StrCat( + return InvalidArgumentStrCat( "Incompatible operand shape at index ", i, ": expected ", ShapeUtil::HumanString(expected_shape), ", got ", - ShapeUtil::HumanString(operand->shape()))); + ShapeUtil::HumanString(operand->shape())); } operand = comp->AddInstruction(HloInstruction::CreateUnary( expected_shape, HloOpcode::kCopy, operand)); @@ -125,10 +124,10 @@ absl::StatusOr HloModuleStitcher::RunImpl( HloInstruction* replacement = call; if (!ShapeUtil::Equal(result_shape, inst->shape())) { if (!ShapeUtil::Compatible(result_shape, inst->shape())) { - return absl::InvalidArgumentError( - absl::StrCat("Incompatible result shape: expected ", - ShapeUtil::HumanString(inst->shape()), ", got ", - ShapeUtil::HumanString(result_shape))); + return InvalidArgumentStrCat("Incompatible result shape: expected ", + ShapeUtil::HumanString(inst->shape()), + ", got ", + ShapeUtil::HumanString(result_shape)); } replacement = comp->AddInstruction(HloInstruction::CreateUnary( inst->shape(), HloOpcode::kCopy, call)); diff --git a/third_party/xla/xla/hlo/transforms/host_offload_legalize.cc b/third_party/xla/xla/hlo/transforms/host_offload_legalize.cc index 2f288a0fecb1a2..7bdd49a2d80d50 100644 --- a/third_party/xla/xla/hlo/transforms/host_offload_legalize.cc +++ b/third_party/xla/xla/hlo/transforms/host_offload_legalize.cc @@ -193,8 +193,7 @@ absl::StatusOr WalkUpMemoryOffload( return InstructionAndIndex(instruction, index); } default: { - return absl::InvalidArgumentError( - absl::StrFormat("Invalid opcode %s", instruction->ToString())); + return InvalidArgument("Invalid opcode %s", instruction->ToString()); } } } @@ -236,14 +235,14 @@ absl::StatusOr> WalkDownMemoryOffload( std::vector callers = call_graph.GetComputationCallers(current_value.instruction->parent()); if (callers.size() != 1 || callers[0]->opcode() != HloOpcode::kWhile) { - return absl::InvalidArgumentError(absl::StrFormat( + return InvalidArgument( "Expected computation \"%s\" to be called only by one caller " "and that caller to be a While. There are %d caller(s): [%s]", current_value.instruction->parent()->name(), callers.size(), absl::StrJoin(callers, ", ", [](std::string* out, const HloInstruction* instr) { absl::StrAppend(out, instr->name()); - }))); + })); } ABSL_RETURN_IF_ERROR(add_gte_for_idx(callers[0], current_value.index)); return results; @@ -324,8 +323,7 @@ absl::StatusOr> WalkDownMemoryOffload( [[fallthrough]]; } default: { - return absl::InvalidArgumentError( - absl::StrFormat("Unrecognized user name: %s", user->name())); + return InvalidArgument("Unrecognized user name: %s", user->name()); } } } @@ -389,10 +387,10 @@ absl::StatusOr> GetNewShapesAfterBitcastReducedRank( const Shape& before_bitcast_shape = bitcast->operand(0)->shape(); if (!(ShapeUtil::IsEffectivelyMostMajorDimension(before_bitcast_shape, 0) && before_bitcast_shape.dimensions(0) == 1)) { - return absl::InternalError( - absl::StrFormat("Only handling bitcasts with majormost dimension " - "of size 1. This bitcast is \"%s\"", - bitcast->ToString())); + return Internal( + "Only handling bitcasts with majormost dimension " + "of size 1. This bitcast is \"%s\"", + bitcast->ToString()); } const Shape new_bitcast_shape = RemoveMajormostDimension(shape_before_copy); VLOG(2) << absl::StreamFormat( @@ -413,10 +411,10 @@ absl::StatusOr> GetNewShapesAfterBitcastIncreasedRank( const Shape& after_bitcast_shape = bitcast->shape(); if (!(ShapeUtil::IsEffectivelyMostMajorDimension(after_bitcast_shape, 0) && after_bitcast_shape.dimensions(0) == 1)) { - return absl::UnimplementedError( - absl::StrFormat("Only handling bitcasts with majormost dimension " - "of size 1. This bitcast is \"%s\"", - bitcast->ToString())); + return Unimplemented( + "Only handling bitcasts with majormost dimension " + "of size 1. This bitcast is \"%s\"", + bitcast->ToString()); } const Shape new_bitcast_shape = AddMajormostDimension(shape_before_copy); VLOG(2) << absl::StreamFormat( @@ -442,12 +440,12 @@ absl::StatusOr> GetNewShapesAfterBitcastSameRank( before_bitcast_shape)) { // Something about the shape other than the layout changes. This is not // supported. - return absl::UnimplementedError(absl::StrFormat( + return Unimplemented( "Only handling bitcasts which change the layout. This bitcast (\"%s\") " "has input shape \"%s\" and output shape \"%s\".", bitcast->name(), bitcast->operand(0)->shape().ToString(/*print_layout=*/true), - bitcast->shape().ToString(/*print_layout=*/true))); + bitcast->shape().ToString(/*print_layout=*/true)); } if (Shape::Equal()(after_bitcast_shape, before_bitcast_shape)) { @@ -508,10 +506,10 @@ absl::StatusOr> GetNewShapesAfterBitcastSameRank( return std::make_pair(new_shape_before_copy, new_shape_after_copy); } } - return absl::UnimplementedError(absl::StrFormat( + return Unimplemented( "Something about this layout changed other than the minor-to-major " "ordering. This is unsuppored. Bitcast: \"%s\"", - bitcast->ToString())); + bitcast->ToString()); } // This function is to be called when we are moving a copy down the graph. The @@ -523,9 +521,9 @@ absl::StatusOr> GetNewShapesAfterBitcast( const Shape& shape_before_copy, const Shape& shape_after_copy) { if (!Shape::Equal().IgnoreLayout()(copy_to_move->operand(0)->shape(), copy_to_move->shape())) { - return absl::InternalError(absl::StrFormat( + return Internal( "Expecting copy to only change instruction's layout. Copy: %s", - copy_to_move->ToString())); + copy_to_move->ToString()); } const Shape& before_bitcast_shape = bitcast->operand(0)->shape(); @@ -555,9 +553,9 @@ absl::StatusOr> GetNewShapesAfterBitcast( } // Dimensionality changes in some other way. - return absl::UnimplementedError(absl::StrFormat( + return Unimplemented( "Bitcast changes dimensionality in an unsupported way. Bitcast: \"%s\"", - bitcast->ToString())); + bitcast->ToString()); } absl::Status MoveCopyDown( @@ -879,7 +877,7 @@ absl::StatusOr ProcessAnnotationForCopyMovement( call_graph->GetComputationCallers(annotation->parent()); if (callers.size() != 1 || callers[0]->opcode() != HloOpcode::kWhile) { - return absl::InvalidArgumentError(absl::StrFormat( + return InvalidArgument( "Expected computation \"%s\" to be called only by one caller " "and that caller to be a While. There are %d caller(s): [%s]", current_value.instruction->parent()->name(), callers.size(), @@ -887,7 +885,7 @@ absl::StatusOr ProcessAnnotationForCopyMovement( callers, ", ", [](std::string* out, const HloInstruction* instr) { absl::StrAppend(out, instr->name()); - }))); + })); } for (int i = 0; i < user->operands().size(); i++) { if (user->operands()[i] == annotation && diff --git a/third_party/xla/xla/hlo/transforms/host_offloader.cc b/third_party/xla/xla/hlo/transforms/host_offloader.cc index 9a7f45cf83a487..7f1c81c48d0403 100644 --- a/third_party/xla/xla/hlo/transforms/host_offloader.cc +++ b/third_party/xla/xla/hlo/transforms/host_offloader.cc @@ -830,10 +830,10 @@ absl::Status HostOffloader::CreateAllocateBufferForDynamicUpdateSlice( // Buffer comes from one parameter. Stop the process. return absl::OkStatus(); } - return absl::InvalidArgumentError( - absl::StrFormat("Entry computation parameter \"%s\" (shape index " - "%s) is not in host memory space.", - instruction->name(), shape_index.ToString())); + return InvalidArgument( + "Entry computation parameter \"%s\" (shape index " + "%s) is not in host memory space.", + instruction->name(), shape_index.ToString()); } // If this is a parameter of a while_body, we also need to find the @@ -872,10 +872,10 @@ absl::Status HostOffloader::CreateAllocateBufferForDynamicUpdateSlice( nested_queue.pop(); if (!host_offload_utils::IsValidDuringPureMemoryOffload( nested_instruction_and_shape.instruction)) { - return absl::InvalidArgumentError(absl::StrFormat( + return InvalidArgument( "Tensor which is moved to host is used by an invalid " "instruction (\"%s\") during while condition body.", - nested_instruction_and_shape.instruction->name())); + nested_instruction_and_shape.instruction->name()); } SetMemorySpace( ShapeUtil::GetMutableSubshape( @@ -1046,10 +1046,10 @@ absl::Status HostOffloader::CreateAllocateBufferForDynamicUpdateSlice( } } if (!found_broadcast) { - return absl::InvalidArgumentError( - absl::StrFormat("DynamicUpdateSlice \"%s\"'s first operand is not the " - "result of a broadcast.", - dynamic_update_slice->name())); + return InvalidArgument( + "DynamicUpdateSlice \"%s\"'s first operand is not the " + "result of a broadcast.", + dynamic_update_slice->name()); } return absl::OkStatus(); } @@ -1138,9 +1138,8 @@ absl::Status ValidateAsyncComputationStructure(HloComputation* computation) { continue; } - return absl::InternalError( - absl::StrCat("Unexpected instruction found in async computation: ", - instr->ToString())); + return InternalStrCat("Unexpected instruction found in async computation: ", + instr->ToString()); } return absl::OkStatus(); diff --git a/third_party/xla/xla/hlo/transforms/simplifiers/all_gather_pad_ds_simplifier_test.cc b/third_party/xla/xla/hlo/transforms/simplifiers/all_gather_pad_ds_simplifier_test.cc index 58a967c6396b7b..ff1b9d44951327 100644 --- a/third_party/xla/xla/hlo/transforms/simplifiers/all_gather_pad_ds_simplifier_test.cc +++ b/third_party/xla/xla/hlo/transforms/simplifiers/all_gather_pad_ds_simplifier_test.cc @@ -60,11 +60,9 @@ class AllGatherPadDsSimplifierTest : public HloHardwareIndependentTestBase { config.set_use_spmd_partitioning(num_partitions > 1); ABSL_ASSIGN_OR_RETURN(auto module, ParseAndReturnVerifiedModule(hlo_module, config)); - auto changed = AllGatherPadDsSimplifier().Run(module.get(), {}); - if (!changed.ok()) { - return changed.status(); - } - EXPECT_EQ(changed.value(), expect_change); + ABSL_ASSIGN_OR_RETURN(bool changed, + AllGatherPadDsSimplifier().Run(module.get(), {})); + EXPECT_EQ(changed, expect_change); LOG(INFO) << "new module: " << module->ToString(); return module; } diff --git a/third_party/xla/xla/hlo/transforms/simplifiers/all_gather_permuted_ds_simplifier_test.cc b/third_party/xla/xla/hlo/transforms/simplifiers/all_gather_permuted_ds_simplifier_test.cc index 68e97976af8584..2a6ae89ab9626a 100644 --- a/third_party/xla/xla/hlo/transforms/simplifiers/all_gather_permuted_ds_simplifier_test.cc +++ b/third_party/xla/xla/hlo/transforms/simplifiers/all_gather_permuted_ds_simplifier_test.cc @@ -53,12 +53,10 @@ class AllGatherPermutedDsSimplifierTest config.set_use_spmd_partitioning(num_partitions > 1); ABSL_ASSIGN_OR_RETURN(std::unique_ptr module, ParseAndReturnVerifiedModule(hlo_module, config)); - absl::StatusOr changed = - AllGatherDynamicSlicePermutedOffsetSimplifier().Run(module.get(), {}); - if (!changed.ok()) { - return changed.status(); - } - EXPECT_EQ(*changed, expect_change); + ABSL_ASSIGN_OR_RETURN( + bool changed, + AllGatherDynamicSlicePermutedOffsetSimplifier().Run(module.get(), {})); + EXPECT_EQ(changed, expect_change); return module; } }; diff --git a/third_party/xla/xla/hlo/transforms/simplifiers/reduce_window_rewriter.cc b/third_party/xla/xla/hlo/transforms/simplifiers/reduce_window_rewriter.cc index d4026335655e7e..c51a3c0d64ceac 100644 --- a/third_party/xla/xla/hlo/transforms/simplifiers/reduce_window_rewriter.cc +++ b/third_party/xla/xla/hlo/transforms/simplifiers/reduce_window_rewriter.cc @@ -154,9 +154,8 @@ static absl::StatusOr ScalarizeComputation( break; default: { if (!inst->IsElementwise()) { - return absl::InvalidArgumentError( - absl::StrCat("Instruction is not elementwise: ", - HloOpcodeString(inst->opcode()))); + return InvalidArgumentStrCat("Instruction is not elementwise: ", + HloOpcodeString(inst->opcode())); } ABSL_ASSIGN_OR_RETURN(Shape shape, get_scalar_shape(inst->shape())); new_inst = builder.AddInstruction( diff --git a/third_party/xla/xla/hlo/transforms/simplifiers/unflatten_call_graph.cc b/third_party/xla/xla/hlo/transforms/simplifiers/unflatten_call_graph.cc index f5430410ecc451..11d0e9e50781b8 100644 --- a/third_party/xla/xla/hlo/transforms/simplifiers/unflatten_call_graph.cc +++ b/third_party/xla/xla/hlo/transforms/simplifiers/unflatten_call_graph.cc @@ -133,18 +133,19 @@ absl::Status UnflattenCallGraph::ValidateComputationHashes( const std::vector& hash_results, const absl::flat_hash_map& hash_to_canonical) { - auto validate_against_canonical = [&](const ComputationHashResult& result) { + auto validate_against_canonical = + [&](const ComputationHashResult& result) -> absl::Status { uint64_t candidate_hash = result.hash; const std::string& candidate_fingerprint = result.fingerprint; const std::string& canonical_fingerprint = hash_to_canonical.at(candidate_hash)->fingerprint; if (candidate_fingerprint != canonical_fingerprint) { - return absl::InternalError( - absl::StrCat("Hash collision detected. Hash: ", candidate_hash, "\n", - "Hashes are equal but fingerprints are different.\n", - "Computation 1:\n", candidate_fingerprint, "\n", - "Computation 2:\n", canonical_fingerprint, "\n")); + return InternalStrCat( + "Hash collision detected. Hash: ", candidate_hash, "\n", + "Hashes are equal but fingerprints are different.\n", + "Computation 1:\n", candidate_fingerprint, "\n", "Computation 2:\n", + canonical_fingerprint, "\n"); } return absl::OkStatus(); }; diff --git a/third_party/xla/xla/util.h b/third_party/xla/xla/util.h index 387599ff99a403..be3577a85f1660 100644 --- a/third_party/xla/xla/util.h +++ b/third_party/xla/xla/util.h @@ -345,8 +345,7 @@ XLA_ERROR_WITH_STRFORMAT_AND_BACKTRACE(Unknown); #define XLA_ERROR_WITH_STRCAT_AND_BACKTRACE_PREFIX(error_type) \ template \ struct error_type##StrCat { \ - absl::Status status; \ - /* NOLINTNEXTLINE(google-explicit-constructor) */ + absl::Status status; #define XLA_ERROR_WITH_STRCAT_AND_BACKTRACE_SUFFIX(error_type) \ /* NOLINTNEXTLINE(google-explicit-constructor) */ \ operator absl::Status() const { return status; } \ @@ -360,8 +359,9 @@ XLA_ERROR_WITH_STRFORMAT_AND_BACKTRACE(Unknown); #if defined(PLATFORM_GOOGLE) #define XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(error_type) \ XLA_ERROR_WITH_STRCAT_AND_BACKTRACE_PREFIX(error_type) \ - error_type##StrCat(Args&&... concat, absl::SourceLocation loc = \ - absl::SourceLocation::current()) \ + explicit error_type##StrCat( \ + Args&&... concat, \ + absl::SourceLocation loc = absl::SourceLocation::current()) \ : status( \ WithLogBacktrace(absl::error_type##Error( \ absl::StrCat(std::forward(concat)...)) \ @@ -370,16 +370,23 @@ XLA_ERROR_WITH_STRFORMAT_AND_BACKTRACE(Unknown); #else #define XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(error_type) \ XLA_ERROR_WITH_STRCAT_AND_BACKTRACE_PREFIX(error_type) \ - error_type##StrCat(Args&&... concat) \ + explicit error_type##StrCat(Args&&... concat) \ : status(WithLogBacktrace(absl::error_type##Error( \ absl::StrCat(std::forward(concat)...)))) {} \ XLA_ERROR_WITH_STRCAT_AND_BACKTRACE_SUFFIX(error_type) #endif -XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(ResourceExhausted); +XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(Aborted); +XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(Cancelled); +XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(DeadlineExceeded); +XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(FailedPrecondition); +XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(Internal); XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(InvalidArgument); +XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(NotFound); +XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(ResourceExhausted); +XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(Unavailable); XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(Unimplemented); -XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(Internal); +XLA_ERROR_WITH_STRCAT_AND_BACKTRACE(Unknown); #undef XLA_ERROR_WITH_STRCAT_AND_BACKTRACE #undef XLA_ERROR_WITH_STRCAT_AND_BACKTRACE_PREFIX diff --git a/third_party/xla/xla/util_test.cc b/third_party/xla/xla/util_test.cc index 5f30032c3e11bd..dc9393e431d628 100644 --- a/third_party/xla/xla/util_test.cc +++ b/third_party/xla/xla/util_test.cc @@ -29,6 +29,7 @@ limitations under the License. #include "absl/algorithm/container.h" #include "absl/container/inlined_vector.h" #include "absl/log/check.h" +#include "absl/status/status.h" #include "absl/strings/match.h" #include "absl/strings/string_view.h" #include "absl/types/span.h" @@ -661,6 +662,25 @@ TEST(UtilTest, ScopedLoggingTimerLazyEvaluation) { EXPECT_EQ(counter, 0); } +TEST(UtilTest, ErrorWithStrCatAndBacktrace) { + absl::Status invalid_arg = InvalidArgumentStrCat("bad arg: ", 42); + EXPECT_EQ(invalid_arg.code(), absl::StatusCode::kInvalidArgument); + EXPECT_THAT(invalid_arg.message(), ::testing::HasSubstr("bad arg: 42")); + + absl::Status internal = InternalStrCat("internal failure: ", "oom"); + EXPECT_EQ(internal.code(), absl::StatusCode::kInternal); + EXPECT_THAT(internal.message(), + ::testing::HasSubstr("internal failure: oom")); + + absl::Status precondition = FailedPreconditionStrCat("state=", 1); + EXPECT_EQ(precondition.code(), absl::StatusCode::kFailedPrecondition); + EXPECT_THAT(precondition.message(), ::testing::HasSubstr("state=1")); + + absl::Status not_found = NotFoundStrCat("missing key: ", "foo"); + EXPECT_EQ(not_found.code(), absl::StatusCode::kNotFound); + EXPECT_THAT(not_found.message(), ::testing::HasSubstr("missing key: foo")); +} + void BM_PackIntN(::testing::benchmark::State& state) { const int bitwidth = state.range(0); const size_t num_elements = state.range(1); From dc599b38461b0742888b4fb9f5ad149cbab4e7ff Mon Sep 17 00:00:00 2001 From: Sam Ginzburg Date: Tue, 29 Sep 2026 16:58:38 -0700 Subject: [PATCH 08/21] [XLA] Add PrecisionConfig algorithms for FP8x3 and FP8x4 BF16 dot emulation This change introduces ALG_DOT_BF16_BF16_FP8X3 and ALG_DOT_BF16_BF16_FP8X4 to PrecisionConfig in xla_data.proto and plumbs them through StableHLO and MHLO attribute translators. PiperOrigin-RevId: 990613954 --- .../xla/third_party/stablehlo/temporary.patch | 3712 +++++++++++++++++ .../tests/fusion_emitter_device_test.cc | 3 + .../hlo_to_mhlo/attribute_importer.cc | 12 + .../hlo_to_mhlo/tests/attributes.hlo | 24 + .../mhlo_to_hlo/attribute_exporter.cc | 8 + .../mhlo_to_hlo/tests/attributes.mlir | 48 + .../mhlo/hlo-legalize-to-stablehlo.mlir | 78 + .../mhlo/stablehlo-legalize-to-hlo.mlir | 82 + third_party/xla/xla/service/algorithm_util.cc | 7 + third_party/xla/xla/xla_data.proto | 10 +- 10 files changed, 3982 insertions(+), 2 deletions(-) diff --git a/third_party/xla/third_party/stablehlo/temporary.patch b/third_party/xla/third_party/stablehlo/temporary.patch index 73906a9c97345e..d452af844efe30 100644 --- a/third_party/xla/third_party/stablehlo/temporary.patch +++ b/third_party/xla/third_party/stablehlo/temporary.patch @@ -548,6 +548,40 @@ diff --ruN a/stablehlo/stablehlo/conversions/tosa/tests/unary.mlir b/stablehlo/s // CHECK: tosa.reshape %arg0, %[[VAR0]] %0 = "stablehlo.reshape"(%arg0) : (tensor<2x3xf32>) -> tensor<6xf32> return %0 : tensor<6xf32> +diff --ruN a/stablehlo/stablehlo/dialect/Base.cpp b/stablehlo/stablehlo/dialect/Base.cpp +--- stablehlo/stablehlo/dialect/Base.cpp ++++ stablehlo/stablehlo/dialect/Base.cpp +@@ -671,6 +671,7 @@ + StringRef f32 = Float32Type::name; + StringRef f64 = Float64Type::name; + StringRef tf32 = FloatTF32Type::name; ++ StringRef f8e4m3fn = Float8E4M3FNType::name; + std::map, + KnownDotAlgorithm> + knownDotAlgorithms{ +@@ -679,6 +680,10 @@ + {{bf16, bf16, bf16, 1}, KnownDotAlgorithm::BF16_BF16_BF16}, + {{bf16, bf16, f32, 1}, KnownDotAlgorithm::BF16_BF16_F32}, + {{bf16, bf16, f32, 3}, KnownDotAlgorithm::BF16_BF16_F32_X3}, ++ {{f8e4m3fn, f8e4m3fn, f32, 3}, ++ KnownDotAlgorithm::F8E4M3FN_F8E4M3FN_F32_X3}, ++ {{f8e4m3fn, f8e4m3fn, f32, 4}, ++ KnownDotAlgorithm::F8E4M3FN_F8E4M3FN_F32_X4}, + {{bf16, bf16, f32, 6}, KnownDotAlgorithm::BF16_BF16_F32_X6}, + {{bf16, bf16, f32, 9}, KnownDotAlgorithm::BF16_BF16_F32_X9}, + {{tf32, tf32, f32, 1}, KnownDotAlgorithm::TF32_TF32_F32}, +diff --ruN a/stablehlo/stablehlo/dialect/Base.h b/stablehlo/stablehlo/dialect/Base.h +--- stablehlo/stablehlo/dialect/Base.h ++++ stablehlo/stablehlo/dialect/Base.h +@@ -264,6 +264,8 @@ + F32_F32_F32 = 11, + F64_F64_F64 = 12, + BF16_BF16_F32_X9 = 13, ++ F8E4M3FN_F8E4M3FN_F32_X3 = 14, ++ F8E4M3FN_F8E4M3FN_F32_X4 = 15, + }; + + FailureOr getKnownDotAlgorithm( diff --ruN a/stablehlo/stablehlo/dialect/ChloBytecode.cpp b/stablehlo/stablehlo/dialect/ChloBytecode.cpp --- stablehlo/stablehlo/dialect/ChloBytecode.cpp +++ stablehlo/stablehlo/dialect/ChloBytecode.cpp @@ -952,6 +986,18 @@ diff --ruN a/stablehlo/stablehlo/dialect/StablehloOps.td b/stablehlo/stablehlo/d }]>, OpBuilder<(ins "::mlir::TypeRange":$resultTypes, "::mlir::ValueRange":$inputs, +diff --ruN a/stablehlo/stablehlo/dialect/Version.h b/stablehlo/stablehlo/dialect/Version.h +--- stablehlo/stablehlo/dialect/Version.h ++++ stablehlo/stablehlo/dialect/Version.h +@@ -38,7 +38,7 @@ + static FailureOr fromString(llvm::StringRef versionRef); + + /// Return a Version representing the current VHLO dialect version. +- static Version getCurrentVersion() { return Version(1, 20, 1); } ++ static Version getCurrentVersion() { return Version(1, 21, 0); } + + /// Return a Version representing the minimum supported VHLO dialect version. + static Version getMinimumVersion() { return Version(0, 9, 0); } diff --ruN a/stablehlo/stablehlo/dialect/VhloBytecode.h b/stablehlo/stablehlo/dialect/VhloBytecode.h --- stablehlo/stablehlo/dialect/VhloBytecode.h +++ stablehlo/stablehlo/dialect/VhloBytecode.h @@ -964,6 +1010,17 @@ diff --ruN a/stablehlo/stablehlo/dialect/VhloBytecode.h b/stablehlo/stablehlo/di } // namespace vhlo } // namespace mlir +diff --ruN a/stablehlo/stablehlo/dialect/VhloDialect.td b/stablehlo/stablehlo/dialect/VhloDialect.td +--- stablehlo/stablehlo/dialect/VhloDialect.td ++++ stablehlo/stablehlo/dialect/VhloDialect.td +@@ -59,6 +59,7 @@ + 1.18.0: Add `result_tilings` attribute to `custom_call` op. + 1.19.0: Add CollectiveReduceOp. + 1.20.0: Add has_dynamic_root to collective_broadcast op. Allow mixed fp8 operands in `convolution` and `dynamic_conv` ops. ++ 1.21.0: Add F8E4M3FN_F8E4M3FN_F32_X3 and F8E4M3FN_F8E4M3FN_F32_X4 dot algorithms. + }]; + + let useDefaultAttributePrinterParser = 0; diff --ruN a/stablehlo/stablehlo/dialect/VhloEnums.td b/stablehlo/stablehlo/dialect/VhloEnums.td --- stablehlo/stablehlo/dialect/VhloEnums.td +++ stablehlo/stablehlo/dialect/VhloEnums.td @@ -975,6 +1032,42 @@ diff --ruN a/stablehlo/stablehlo/dialect/VhloEnums.td b/stablehlo/stablehlo/dial let extraClassDeclaration = [{ mlir::vhlo::Version getMinVersion() { return mlir::vhlo::Version(}] # !subst(".", ", ", minVersion) # [{); +diff --ruN a/stablehlo/stablehlo/dialect/VhloOps.cpp b/stablehlo/stablehlo/dialect/VhloOps.cpp +--- stablehlo/stablehlo/dialect/VhloOps.cpp ++++ stablehlo/stablehlo/dialect/VhloOps.cpp +@@ -382,6 +382,20 @@ + return success(); + } + ++LogicalResult verifyConstraint_1_21_0(mlir::Operation* op, ++ Version targetVersion) { ++ auto dotGeneralOp = cast(op); ++ if (targetVersion < Version(1, 21, 0)) { ++ auto lhsType = dyn_cast(dotGeneralOp.getLhsPrecisionType()); ++ auto rhsType = dyn_cast(dotGeneralOp.getRhsPrecisionType()); ++ if ((lhsType && isa(lhsType.getValue())) || ++ (rhsType && isa(rhsType.getValue()))) { ++ return failure(); ++ } ++ } ++ return success(); ++} ++ + } // namespace + + LogicalResult AllReduceOpV1::validateConstraint(mlir::Operation* op, +@@ -394,6 +408,11 @@ + return verifyConstraint_1_20_0(op, targetVersion); + } + ++LogicalResult DotGeneralOpV2::validateConstraint(mlir::Operation* op, ++ Version targetVersion) { ++ return verifyConstraint_1_21_0(op, targetVersion); ++} ++ + LogicalResult DynamicConvOpV2::validateConstraint(mlir::Operation* op, + Version targetVersion) { + return verifyConstraint_1_20_0(op, targetVersion); diff --ruN a/stablehlo/stablehlo/dialect/VhloOps.h b/stablehlo/stablehlo/dialect/VhloOps.h --- stablehlo/stablehlo/dialect/VhloOps.h +++ stablehlo/stablehlo/dialect/VhloOps.h @@ -1004,6 +1097,19 @@ diff --ruN a/stablehlo/stablehlo/dialect/VhloOps.h b/stablehlo/stablehlo/dialect private: // Adds VHLO types to this dialect. +diff --ruN a/stablehlo/stablehlo/dialect/VhloOps.td b/stablehlo/stablehlo/dialect/VhloOps.td +--- stablehlo/stablehlo/dialect/VhloOps.td ++++ stablehlo/stablehlo/dialect/VhloOps.td +@@ -512,7 +512,8 @@ + let results = (outs VHLO_AnyType:$result); + } + +-def VHLO_DotGeneralOpV2 : VHLO_Op<"dot_general_v2", "1.6.0", "current"> { ++def VHLO_DotGeneralOpV2 : VHLO_Op<"dot_general_v2", "1.6.0", "current", ++ [DeclareOpInterfaceMethods]> { + let arguments = (ins + VHLO_AnyType:$lhs, + VHLO_AnyType:$rhs, diff --ruN a/stablehlo/stablehlo/integrations/c/ChloAttributes.cpp b/stablehlo/stablehlo/integrations/c/ChloAttributes.cpp --- stablehlo/stablehlo/integrations/c/ChloAttributes.cpp +++ stablehlo/stablehlo/integrations/c/ChloAttributes.cpp @@ -2132,6 +2238,50 @@ diff --ruN a/stablehlo/stablehlo/tests/ops_chlo.mlir b/stablehlo/stablehlo/tests // CHECK-LABEL: func @mulhi_i32 func.func @mulhi_i32(%arg0: tensor<4xi32>, %arg1: tensor<4xi32>) -> tensor<4xi32> { %0 = "chlo.mulhi"(%arg0, %arg1) : (tensor<4xi32>, tensor<4xi32>) -> tensor<4xi32> +diff --ruN a/stablehlo/stablehlo/tests/ops_dot_general_algorithms.mlir b/stablehlo/stablehlo/tests/ops_dot_general_algorithms.mlir +--- stablehlo/stablehlo/tests/ops_dot_general_algorithms.mlir ++++ stablehlo/stablehlo/tests/ops_dot_general_algorithms.mlir +@@ -119,6 +119,40 @@ + }> : (tensor<2x2x2xbf16>, tensor<2x2x2xbf16>) -> tensor<2x2x2xbf16> return %0 : tensor<2x2x2xbf16> + } + ++// CHECK-LABEL: func @dot_algorithm_f8e4m3fn_f8e4m3fn_f32_x3 ++func.func @dot_algorithm_f8e4m3fn_f8e4m3fn_f32_x3(%arg0: tensor<2x2x2xbf16>, %arg1: tensor<2x2x2xbf16>) -> tensor<2x2x2xf32> { ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) <{ ++ dot_dimension_numbers = #stablehlo.dot, ++ precision_config = [#stablehlo, #stablehlo], ++ algorithm = #stablehlo.dot_algorithm< ++ lhs_precision_type = f8E4M3FN, ++ rhs_precision_type = f8E4M3FN, ++ accumulation_type = f32, ++ lhs_component_count = 1, ++ rhs_component_count = 1, ++ num_primitive_operations = 3, ++ allow_imprecise_accumulation = false ++ > ++ }> : (tensor<2x2x2xbf16>, tensor<2x2x2xbf16>) -> tensor<2x2x2xf32> return %0 : tensor<2x2x2xf32> ++} ++ ++// CHECK-LABEL: func @dot_algorithm_f8e4m3fn_f8e4m3fn_f32_x4 ++func.func @dot_algorithm_f8e4m3fn_f8e4m3fn_f32_x4(%arg0: tensor<2x2x2xbf16>, %arg1: tensor<2x2x2xbf16>) -> tensor<2x2x2xf32> { ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) <{ ++ dot_dimension_numbers = #stablehlo.dot, ++ precision_config = [#stablehlo, #stablehlo], ++ algorithm = #stablehlo.dot_algorithm< ++ lhs_precision_type = f8E4M3FN, ++ rhs_precision_type = f8E4M3FN, ++ accumulation_type = f32, ++ lhs_component_count = 1, ++ rhs_component_count = 1, ++ num_primitive_operations = 4, ++ allow_imprecise_accumulation = false ++ > ++ }> : (tensor<2x2x2xbf16>, tensor<2x2x2xbf16>) -> tensor<2x2x2xf32> return %0 : tensor<2x2x2xf32> ++} ++ + // CHECK-LABEL: func @dot_algorithm_bf16_bf16_f32_x6 + func.func @dot_algorithm_bf16_bf16_f32_x6(%arg0: tensor<2x2x2xbf16>, %arg1: tensor<2x2x2xbf16>) -> tensor<2x2x2xbf16> { + %0 = "stablehlo.dot_general"(%arg0, %arg1) <{ diff --ruN a/stablehlo/stablehlo/tests/transforms/mesh_axes_replica_group_compatibility.mlir b/stablehlo/stablehlo/tests/transforms/mesh_axes_replica_group_compatibility.mlir --- stablehlo/stablehlo/tests/transforms/mesh_axes_replica_group_compatibility.mlir +++ stablehlo/stablehlo/tests/transforms/mesh_axes_replica_group_compatibility.mlir @@ -2194,6 +2344,3568 @@ diff --ruN a/stablehlo/stablehlo/tests/transforms/stablehlo_refine_shapes.mlir b %2 = stablehlo.dynamic_iota %1, dim = 0 : (tensor<1xi64>) -> tensor return %2 : tensor } +diff --ruN a/stablehlo/stablehlo/tests/vhlo/stablehlo_legalize_to_vhlo.1_21_0.mlir b/stablehlo/stablehlo/tests/vhlo/stablehlo_legalize_to_vhlo.1_21_0.mlir +--- stablehlo/stablehlo/tests/vhlo/stablehlo_legalize_to_vhlo.1_21_0.mlir ++++ stablehlo/stablehlo/tests/vhlo/stablehlo_legalize_to_vhlo.1_21_0.mlir +@@ -0,0 +1,3425 @@ ++// RUN: stablehlo-opt --mlir-print-op-generic %s.bc | FileCheck %s ++// RUN: stablehlo-translate --deserialize %s.bc | stablehlo-translate --serialize --target=1.21.0 | stablehlo-opt --mlir-print-op-generic | FileCheck %s ++// RUN: stablehlo-translate --deserialize %s.bc | stablehlo-opt > %t.0 ++// RUN: stablehlo-opt --strip-debuginfo %s > %t.1 ++// RUN: diff %t.0 %t.1 ++// RUN: stablehlo-translate --serialize --target=1.21.0 --strip-debuginfo %s > %t.2 ++// RUN: diff %s.bc %t.2 ++// RUN: %if asserts %{ stablehlo-opt --stablehlo-legalize-to-vhlo -emit-bytecode -debug-only=vhlo-bytecode %s 2>&1 | FileCheck --check-prefix=CHECK-WARN %s %} ++// RUN: %if asserts %{ stablehlo-opt --stablehlo-legalize-to-vhlo -emit-bytecode %s | stablehlo-opt -debug-only=vhlo-bytecode 2>&1 | FileCheck --check-prefix=CHECK-WARN %s %} ++ ++// CHECK-WARN-NOT: Not Implemented ++ ++func.func private @mesh() ++ ++// ============ ATTRIBUTES ============ ++ ++// CHECK-LABEL: "attr_comparison_direction_eq" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_direction_eq(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ // CHECK: comparison_direction = #vhlo ++ comparison_direction = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_direction_ne" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_direction_ne(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ // CHECK: comparison_direction = #vhlo ++ comparison_direction = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_direction_ge" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_direction_ge(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ // CHECK: comparison_direction = #vhlo ++ comparison_direction = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_direction_gt" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_direction_gt(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ // CHECK: comparison_direction = #vhlo ++ comparison_direction = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_direction_le" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_direction_le(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ // CHECK: comparison_direction = #vhlo ++ comparison_direction = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_direction_lt" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_direction_lt(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ // CHECK: comparison_direction = #vhlo ++ comparison_direction = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_type_notype" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_type_notype(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ comparison_direction = #stablehlo ++ // CHECK: compare_type = #vhlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_type_float" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_type_float(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ comparison_direction = #stablehlo, ++ // CHECK: compare_type = #vhlo, ++ compare_type = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_type_totalorder" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_type_totalorder(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ comparison_direction = #stablehlo, ++ // CHECK: compare_type = #vhlo, ++ compare_type = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_type_signed" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_type_signed(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ comparison_direction = #stablehlo, ++ // CHECK: compare_type = #vhlo, ++ compare_type = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_comparison_type_unsigned" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_comparison_type_unsigned(%arg0: tensor, %arg1: tensor) -> tensor { ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ comparison_direction = #stablehlo, ++ // CHECK: compare_type = #vhlo, ++ compare_type = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// ConvDimensionNumbers aka #stablehlo.conv is covered below. ++ ++// CHECK-LABEL: "attr_custom_call_api_version_unspecified" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_custom_call_api_version_unspecified(%arg0: tensor) -> tensor { ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "foo", ++ // CHECK: api_version = #vhlo ++ api_version = 0 : i32 ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_custom_call_api_version_original" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_custom_call_api_version_original(%arg0: tensor) -> tensor { ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "foo", ++ // CHECK: api_version = #vhlo ++ api_version = 1 : i32 ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_custom_call_api_version_status_returning" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_custom_call_api_version_status_returning(%arg0: tensor) -> tensor { ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "foo", ++ // CHECK: api_version = #vhlo ++ api_version = 2 : i32 ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_custom_call_api_version_status_returning_unified" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_custom_call_api_version_status_returning_unified(%arg0: tensor) -> tensor { ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "foo", ++ // CHECK: api_version = #vhlo ++ api_version = 3 : i32 ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_dict" ++// CHECK: #vhlo.dict_v1<{#vhlo.string_v1<"attr1"> = #vhlo.integer_v1<1 : i32>, #vhlo.string_v1<"attr2"> = #vhlo.integer_v1<2 : i32>} ++func.func @attr_dict() attributes {stablehlo.attr = {attr1 = 1 : i32, attr2 = 2 : i32}} { ++ return ++} ++ ++// CHECK-LABEL: "attr_custom_call_api_version_typed_ffi" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++// CHECK: api_version = #vhlo ++// CHECK-SAME: backend_config = #vhlo.dict_v1<{#vhlo.string_v1<"bar"> = #vhlo.integer_v1<42 : i32>}> ++func.func @attr_custom_call_api_version_typed_ffi(%arg0: tensor) -> tensor { ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "foo", ++ backend_config= {bar = 42 : i32}, ++ api_version = 4 : i32 ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++ ++// CHECK-LABEL: "attr_custom_call_api_version_typed_ffi_no_backend_config" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++// CHECK: api_version = #vhlo ++// CHECK-SAME: backend_config = #vhlo.dict_v1<{}> ++func.func @attr_custom_call_api_version_typed_ffi_no_backend_config(%arg0: tensor) -> tensor { ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "foo", ++ api_version = 4 : i32 ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// DotDimensionNumbers aka #stablehlo.dot is covered below. ++ ++// CHECK-LABEL: "attr_fft_type_fft" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_fft_type_fft(%arg0: tensor<16xcomplex>) -> tensor<16xcomplex> { ++ %0 = "stablehlo.fft"(%arg0) { ++ // CHECK: fft_type = #vhlo ++ fft_type = #stablehlo, ++ fft_length = array ++ } : (tensor<16xcomplex>) -> tensor<16xcomplex> ++ func.return %0 : tensor<16xcomplex> ++} ++ ++// CHECK-LABEL: "attr_fft_type_ifft" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_fft_type_ifft(%arg0: tensor<16xcomplex>) -> tensor<16xcomplex> { ++ %0 = "stablehlo.fft"(%arg0) { ++ // CHECK: fft_type = #vhlo ++ fft_type = #stablehlo, ++ fft_length = array ++ } : (tensor<16xcomplex>) -> tensor<16xcomplex> ++ func.return %0 : tensor<16xcomplex> ++} ++ ++// CHECK-LABEL: "attr_fft_type_rfft" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_fft_type_rfft(%arg0: tensor<16xf32>) -> tensor<9xcomplex> { ++ %0 = "stablehlo.fft"(%arg0) { ++ // CHECK: fft_type = #vhlo ++ fft_type = #stablehlo, ++ fft_length = array ++ } : (tensor<16xf32>) -> tensor<9xcomplex> ++ func.return %0 : tensor<9xcomplex> ++} ++ ++// CHECK-LABEL: "attr_fft_type_irfft" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_fft_type_irfft(%arg0: tensor<9xcomplex>) -> tensor<16xf32> { ++ %0 = "stablehlo.fft"(%arg0) { ++ // CHECK: fft_type = #vhlo ++ fft_type = #stablehlo, ++ fft_length = array ++ } : (tensor<9xcomplex>) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "exponential_HIGHEST" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}} ++func.func @exponential_HIGHEST(%arg0: tensor<8x16xf32>) -> tensor<8x16xf32> { ++ %0 = "stablehlo.exponential"(%arg0) { ++ // CHECK: result_accuracy = #vhlo.result_accuracy_v1> ++ result_accuracy = #stablehlo.result_accuracy> ++ } : (tensor<8x16xf32>) -> tensor<8x16xf32> ++ func.return %0 : tensor<8x16xf32> ++} ++ ++// CHECK-LABEL: "exponential_TOLERANCE" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}} ++func.func @exponential_TOLERANCE(%arg0: tensor<8x16xf32>) -> tensor<8x16xf32> { ++ %0 = "stablehlo.exponential"(%arg0) { ++ // CHECK: result_accuracy = #vhlo.result_accuracy_v1> ++ result_accuracy = #stablehlo.result_accuracy> ++ } : (tensor<8x16xf32>) -> tensor<8x16xf32> ++ func.return %0 : tensor<8x16xf32> ++} ++ ++// CHECK-LABEL: "exponential_DEFAULT" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}} ++func.func @exponential_DEFAULT(%arg0: tensor<8x16xf32>) -> tensor<8x16xf32> { ++ %0 = "stablehlo.exponential"(%arg0) { ++ // CHECK: result_accuracy = #vhlo.result_accuracy_v1> ++ result_accuracy = #stablehlo.result_accuracy> ++ } : (tensor<8x16xf32>) -> tensor<8x16xf32> ++ func.return %0 : tensor<8x16xf32> ++} ++ ++// GatherDimensionNumbers aka #stablehlo.gather is covered below. ++ ++// CHECK-LABEL: "attr_precision_config_default" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_precision_config_default(%arg0: tensor<8x16xf32>, %arg1: tensor<16x8xf32>) -> tensor<8x8xf32> { ++ %0 = "stablehlo.dot"(%arg0, %arg1) { ++ // CHECK: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]> ++ } : (tensor<8x16xf32>, tensor<16x8xf32>) -> tensor<8x8xf32> ++ func.return %0 : tensor<8x8xf32> ++} ++ ++// CHECK-LABEL: "attr_precision_config_high" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_precision_config_high(%arg0: tensor<8x16xf32>, %arg1: tensor<16x8xf32>) -> tensor<8x8xf32> { ++ %0 = "stablehlo.dot"(%arg0, %arg1) { ++ // CHECK: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]> ++ precision_config = [#stablehlo, #stablehlo] ++ } : (tensor<8x16xf32>, tensor<16x8xf32>) -> tensor<8x8xf32> ++ func.return %0 : tensor<8x8xf32> ++} ++ ++// CHECK-LABEL: "attr_precision_config_highest" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_precision_config_highest(%arg0: tensor<8x16xf32>, %arg1: tensor<16x8xf32>) -> tensor<8x8xf32> { ++ %0 = "stablehlo.dot"(%arg0, %arg1) { ++ // CHECK: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]> ++ precision_config = [#stablehlo, #stablehlo] ++ } : (tensor<8x16xf32>, tensor<16x8xf32>) -> tensor<8x8xf32> ++ func.return %0 : tensor<8x8xf32> ++} ++ ++// CHECK-LABEL: "attr_rng_algorithm_default" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_rng_algorithm_default(%arg0: tensor) -> (tensor, tensor) { ++ %0:2 = "stablehlo.rng_bit_generator"(%arg0) { ++ // CHECK: rng_algorithm = #vhlo ++ rng_algorithm = #stablehlo ++ } : (tensor) -> (tensor, tensor) ++ func.return %0#0, %0#1 : tensor, tensor ++} ++ ++// CHECK-LABEL: "attr_rng_algorithm_three_fry" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_rng_algorithm_three_fry(%arg0: tensor) -> (tensor, tensor) { ++ %0:2 = "stablehlo.rng_bit_generator"(%arg0) { ++ // CHECK: rng_algorithm = #vhlo ++ rng_algorithm = #stablehlo ++ } : (tensor) -> (tensor, tensor) ++ func.return %0#0, %0#1 : tensor, tensor ++} ++ ++// CHECK-LABEL: "attr_rng_algorithm_philox" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_rng_algorithm_philox(%arg0: tensor) -> (tensor, tensor) { ++ %0:2 = "stablehlo.rng_bit_generator"(%arg0) { ++ // CHECK: rng_algorithm = #vhlo ++ rng_algorithm = #stablehlo ++ } : (tensor) -> (tensor, tensor) ++ func.return %0#0, %0#1 : tensor, tensor ++} ++ ++// CHECK-LABEL: "attr_rng_distribution_uniform" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @attr_rng_distribution_uniform(%arg0: tensor, %arg1: tensor, %arg2: tensor<0xindex>) -> tensor { ++ %0 = "stablehlo.rng"(%arg0, %arg1, %arg2) { ++ // CHECK: rng_distribution = #vhlo ++ rng_distribution = #stablehlo ++ } : (tensor, tensor, tensor<0xindex>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_rng_distribution_normal" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @attr_rng_distribution_normal(%arg0: tensor, %arg1: tensor, %arg2: tensor<0xindex>) -> tensor { ++ %0 = "stablehlo.rng"(%arg0, %arg1, %arg2) { ++ // CHECK: rng_distribution = #vhlo ++ rng_distribution = #stablehlo ++ } : (tensor, tensor, tensor<0xindex>) -> tensor ++ func.return %0 : tensor ++} ++ ++// ScatterDimensionNumbers aka #stablehlo.scatter is covered below. ++ ++// CHECK-LABEL: "attr_transpose_no_transpose" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_transpose_no_transpose(%arg0: tensor<16x16xf32>, %arg1: tensor<16x16xf32>) -> tensor<16x16xf32> { ++ %0 = "stablehlo.triangular_solve"(%arg0, %arg1) { ++ left_side = true, ++ lower = true, ++ unit_diagonal = true, ++ // transpose_a = #vhlo, ++ transpose_a = #stablehlo ++ } : (tensor<16x16xf32>, tensor<16x16xf32>) -> tensor<16x16xf32> ++ func.return %0 : tensor<16x16xf32> ++} ++ ++// CHECK-LABEL: "attr_transpose_transpose" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_transpose_transpose(%arg0: tensor<16x16xf32>, %arg1: tensor<16x16xf32>) -> tensor<16x16xf32> { ++ %0 = "stablehlo.triangular_solve"(%arg0, %arg1) { ++ left_side = true, ++ lower = true, ++ unit_diagonal = true, ++ // transpose_a = #vhlo, ++ transpose_a = #stablehlo ++ } : (tensor<16x16xf32>, tensor<16x16xf32>) -> tensor<16x16xf32> ++ func.return %0 : tensor<16x16xf32> ++} ++ ++// CHECK-LABEL: "attr_transpose_adjoint" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @attr_transpose_adjoint(%arg0: tensor<16x16xf32>, %arg1: tensor<16x16xf32>) -> tensor<16x16xf32> { ++ %0 = "stablehlo.triangular_solve"(%arg0, %arg1) { ++ left_side = true, ++ lower = true, ++ unit_diagonal = true, ++ // transpose_a = #vhlo, ++ transpose_a = #stablehlo ++ } : (tensor<16x16xf32>, tensor<16x16xf32>) -> tensor<16x16xf32> ++ func.return %0 : tensor<16x16xf32> ++} ++ ++// TypeExtensionsAttr aka #stablehlo.type_extensions is covered below. ++ ++// CHECK-LABEL: "attr_type_extensions_bounds" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @attr_type_extensions_bounds(%arg0: tensor>) -> tensor> { ++ // CHECK: "vhlo.return_v1"(%[[ARG0]]) : (!vhlo.tensor_v1>) -> () ++ func.return %arg0 : tensor> ++} ++ ++// CHECK-LABEL: "attr_frontend_attributes" ++func.func @attr_frontend_attributes(%arg0: tensor) -> tensor { ++ // CHECK: some.unregistered_attr ++ %1 = stablehlo.cosine %arg0 {some.unregistered_attr = 1 : i32} : tensor ++ return %1 : tensor ++} ++ ++// Builtin attriubute tests ++ ++// CHECK-LABEL: "byte_packed_boolean" ++func.func @byte_packed_boolean() -> (tensor<8xi1>, tensor<8xi1>, tensor<4xi1>, tensor<4xi1>, tensor<16xi1>) { ++ // CHECK: #vhlo.tensor_v1 : tensor<8xi1>> ++ // CHECK-NEXT: #vhlo.tensor_v1 : tensor<4xi1>> ++ // CHECK-NEXT: #vhlo.tensor_v1 : tensor<4xi1>> ++ // CHECK-NEXT: #vhlo.tensor_v1 : tensor<16xi1>> ++ %c = stablehlo.constant dense<[true, false, false, false, false, false, false, false]> : tensor<8xi1> ++ %c_0 = stablehlo.constant dense : tensor<8xi1> ++ %c_1 = stablehlo.constant dense<[true, false, false, false]> : tensor<4xi1> ++ %c_2 = stablehlo.constant dense : tensor<4xi1> ++ %c_3 = stablehlo.constant dense<[true, false, false, false, false, false, false, false, true, false, false, false, false, false, false, false]> : tensor<16xi1> ++ return %c, %c_0, %c_1, %c_2, %c_3 : tensor<8xi1>, tensor<8xi1>, tensor<4xi1>, tensor<4xi1>, tensor<16xi1> ++} ++ ++// ============ DEFAULTS ============ ++ ++// CHECK-LABEL: "default_all_gather" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_all_gather(%arg0: tensor<16x8xf32>) -> tensor<16x16xf32> { ++ // CHECK: "vhlo.all_gather_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: all_gather_dim = #vhlo.integer_v1<1 : i64> ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<2x1xi64>>, ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x16x!vhlo.f32_v1> ++ %0 = "stablehlo.all_gather"(%arg0) { ++ all_gather_dim = 1 : i64, ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64> ++ } : (tensor<16x8xf32>) -> tensor<16x16xf32> ++ func.return %0 : tensor<16x16xf32> ++} ++ ++// CHECK-LABEL: "default_all_gather_variadic" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_all_gather_variadic(%arg0: tensor<16x8xf32>, %arg1: tensor<16x8xf32>) -> (tensor<16x16xf32>, tensor<16x16xf32>) { ++ %0:2 = "stablehlo.all_gather"(%arg0, %arg1) { ++ all_gather_dim = 1 : i64, ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64> ++ } : (tensor<16x8xf32>, tensor<16x8xf32>) -> (tensor<16x16xf32>, tensor<16x16xf32>) ++ func.return %0#0, %0#1 : tensor<16x16xf32>, tensor<16x16xf32> ++} ++ ++// CHECK-LABEL: "default_all_reduce" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_all_reduce(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.all_reduce_v2"(%[[ARG0]]) ++ // CHECK-SAME: <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<2x1xi64>>, ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.add_v1"(%[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ ++ %0 = "stablehlo.all_reduce"(%arg0) ({ ++ ^bb0(%arg1: tensor, %arg2: tensor): ++ %1 = "stablehlo.add"(%arg1, %arg2) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64> ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "attr_replica_group_mesh_axes" ++func.func @attr_replica_group_mesh_axes(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.all_reduce_v2" ++ // CHECK-SAME: replica_groups = #vhlo.replica_group_mesh_axes_v1< ++ // CHECK-SAME: mesh = #vhlo.string_v1<"mesh">, ++ // CHECK-SAME: axes = #vhlo.array_v1<[#vhlo.axis_ref_v1, sub_axis_info = #vhlo.sub_axis_info_v1>, #vhlo.axis_ref_v1>]> ++ // CHECK-SAME: > ++ %0 = "stablehlo.all_reduce"(%arg0) ({ ++ ^bb0(%arg1: tensor, %arg2: tensor): ++ %1 = "stablehlo.add"(%arg1, %arg2) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ replica_groups = #stablehlo.replica_group_mesh_axes< ++ mesh = @mesh, ++ axes = [ ++ #stablehlo.axis_ref, ++ #stablehlo.axis_ref ++ ] ++ > ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++ ++ ++// CHECK-LABEL: "default_all_to_all" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_all_to_all(%arg0: tensor<4x16xf32>) -> tensor<16x4xf32> { ++ // CHECK: "vhlo.all_to_all_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: concat_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<1x4xi64>>, ++ // CHECK-SAME: split_count = #vhlo.integer_v1<4 : i64> ++ // CHECK-SAME: split_dimension = #vhlo.integer_v1<1 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<4x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x4x!vhlo.f32_v1> ++ %0 = "stablehlo.all_to_all"(%arg0) { ++ split_dimension = 1 : i64, ++ concat_dimension = 0 : i64, ++ split_count = 4 : i64, ++ replica_groups = dense<[[0, 1, 2, 3]]> : tensor<1x4xi64> ++ } : (tensor<4x16xf32>) -> tensor<16x4xf32> ++ func.return %0 : tensor<16x4xf32> ++} ++ ++// CHECK-LABEL: "default_all_to_all_variadic" ++func.func @default_all_to_all_variadic(%arg0: tensor<4x16xf32>, %arg1: tensor<5x16xf32>) -> (tensor<16x4xf32>, tensor<20x4xf32>) { ++ %0:2 = "stablehlo.all_to_all"(%arg0, %arg1) { ++ split_dimension = 1 : i64, ++ concat_dimension = 0 : i64, ++ split_count = 4 : i64, ++ replica_groups = dense<[[0, 1, 2, 3]]> : tensor<1x4xi64>, ++ channel_handle = #stablehlo.channel_handle ++ } : (tensor<4x16xf32>, tensor<5x16xf32>) -> (tensor<16x4xf32>, tensor<20x4xf32>) ++ func.return %0#0, %0#1 : tensor<16x4xf32>, tensor<20x4xf32> ++} ++ ++// CHECK-LABEL: "default_cholesky" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_cholesky(%arg0: tensor<1x16x16xf32>) -> tensor<1x16x16xf32> { ++ // CHECK: "vhlo.cholesky_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: lower = #vhlo.bool_v1 ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<1x16x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<1x16x16x!vhlo.f32_v1> ++ %0 = "stablehlo.cholesky"(%arg0) : (tensor<1x16x16xf32>) -> tensor<1x16x16xf32> ++ func.return %0 : tensor<1x16x16xf32> ++} ++ ++// CHECK-LABEL: "default_collective_permute" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_collective_permute(%arg0: tensor<16x8xf32>) -> tensor<16x8xf32> { ++ // CHECK: "vhlo.collective_permute_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: source_target_pairs = #vhlo.tensor_v1 : tensor<3x2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x8x!vhlo.f32_v1> ++ %0 = "stablehlo.collective_permute"(%arg0) { ++ source_target_pairs = dense<[[0, 1], [1, 2], [2, 3]]> : tensor<3x2xi64> ++ } : (tensor<16x8xf32>) -> tensor<16x8xf32> ++ func.return %0 : tensor<16x8xf32> ++} ++ ++// CHECK-LABEL: "default_collective_broadcast" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_collective_broadcast(%arg0: tensor<16x8xf32>) -> tensor<16x8xf32> { ++ // CHECK: "vhlo.collective_broadcast_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: has_dynamic_root = #vhlo.bool_v1, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<1x2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x8x!vhlo.f32_v1> ++ %0 = "stablehlo.collective_broadcast"(%arg0) { ++ replica_groups = dense<[[0, 1]]> : tensor<1x2xi64> ++ } : (tensor<16x8xf32>) -> tensor<16x8xf32> ++ func.return %0 : tensor<16x8xf32> ++} ++ ++// CHECK-LABEL: "default_collective_reduce" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_collective_reduce(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.collective_reduce_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: has_dynamic_root = #vhlo.bool_v1, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<1x2xi64>>, ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.add_v1"(%[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.collective_reduce"(%arg0) ({ ++ ^bb0(%arg1: tensor, %arg2: tensor): ++ %1 = "stablehlo.add"(%arg1, %arg2) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ replica_groups = dense<[[0, 1]]> : tensor<1x2xi64> ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "default_compare" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @default_compare(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.compare_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: compare_type = #vhlo, ++ // CHECK-SAME: comparison_direction = #vhlo ++ // CHECK-SAME: }> : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ comparison_direction = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "default_composite" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_composite(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.composite_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: composite_attributes = #vhlo.dict_v1<{}> ++ // CHECK-SAME: decomposition = #vhlo.string_v1<"composite_target"> ++ // CHECK-SAME: name = #vhlo.string_v1<"stablehlo.composite_target"> ++ // CHECK-SAME: version = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.composite"(%arg0) { ++ name = "stablehlo.composite_target", ++ decomposition = @composite_target ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "default_convolution" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @default_convolution(%arg0: tensor<1x8x8x207xf32>, %arg1: tensor<3x3x207x16xf32>) -> tensor<1x6x6x16xf32> { ++ // CHECK: "vhlo.convolution_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: batch_group_count = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: feature_group_count = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: input_batch_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: input_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: input_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: kernel_input_feature_dimension = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: kernel_output_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: kernel_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: lhs_dilation = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: output_batch_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: output_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: output_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: padding = #vhlo.tensor_v1 : tensor<2x2xi64>>, ++ // CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++ // CHECK-SAME: rhs_dilation = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: window_reversal = #vhlo.tensor_v1 : tensor<2xi1>>, ++ // CHECK-SAME: window_strides = #vhlo.tensor_v1 : tensor<2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<1x8x8x207x!vhlo.f32_v1>, !vhlo.tensor_v1<3x3x207x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<1x6x6x16x!vhlo.f32_v1> ++ %0 = "stablehlo.convolution"(%arg0, %arg1) { ++ dimension_numbers = #stablehlo.conv<[b, 0, 1, f]x[0, 1, i, o]->[b, 0, 1, f]>, ++ feature_group_count = 1 : i64, ++ batch_group_count = 1 : i64 ++ } : (tensor<1x8x8x207xf32>, tensor<3x3x207x16xf32>) -> tensor<1x6x6x16xf32> ++ func.return %0 : tensor<1x6x6x16xf32> ++} ++ ++// CHECK-LABEL: "default_custom_call" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_custom_call(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.custom_call_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: api_version = #vhlo, ++ // CHECK-SAME: backend_config = #vhlo.string_v1<"">, ++ // CHECK-SAME: call_target_name = #vhlo.string_v1<"foo">, ++ // CHECK-SAME: called_computations = #vhlo.array_v1<[]>, ++ // CHECK-SAME: has_side_effect = #vhlo.bool_v1, ++ // CHECK-SAME: operand_layouts = #vhlo.array_v1<[]>, ++ // CHECK-SAME: output_operand_aliases = #vhlo.array_v1<[]>, ++ // CHECK-SAME: result_layouts = #vhlo.array_v1<[]>, ++ // CHECK-SAME: result_tilings = #vhlo.array_v1<[]> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "foo" ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "default_dot_general" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @default_dot_general(%arg0: tensor<8x8x16xf32>, %arg1: tensor<8x16x8xf32>) -> tensor<8x8x8xf32> { ++ // CHECK: "vhlo.dot_general_v2"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: accumulation_type = #vhlo.type_v1, ++ // CHECK-SAME: allow_imprecise_accumulation = #vhlo.type_v1, ++ // CHECK-SAME: lhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: lhs_component_count = #vhlo.type_v1, ++ // CHECK-SAME: lhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: lhs_precision_type = #vhlo.type_v1, ++ // CHECK-SAME: num_primitive_operations = #vhlo.type_v1, ++ // CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++ // CHECK-SAME: rhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: rhs_component_count = #vhlo.type_v1, ++ // CHECK-SAME: rhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: rhs_precision_type = #vhlo.type_v1 ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<8x8x16x!vhlo.f32_v1>, !vhlo.tensor_v1<8x16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<8x8x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) { ++ dot_dimension_numbers = #stablehlo.dot< ++ lhs_batching_dimensions = [0], ++ lhs_contracting_dimensions = [2], ++ rhs_batching_dimensions = [0], ++ rhs_contracting_dimensions = [1] ++ > ++ } : (tensor<8x8x16xf32>, tensor<8x16x8xf32>) -> tensor<8x8x8xf32> ++ func.return %0 : tensor<8x8x8xf32> ++} ++ ++// CHECK-LABEL: "dot_general_algorithm" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @dot_general_algorithm(%arg0: tensor<8x8x16xf32>, %arg1: tensor<8x16x8xf32>) -> tensor<8x8x8xf32> { ++// CHECK: "vhlo.dot_general_v2"(%[[ARG0]], %[[ARG1]]) <{ ++// CHECK-SAME: accumulation_type = #vhlo.type_v1, ++// CHECK-SAME: allow_imprecise_accumulation = #vhlo.bool_v1, ++// CHECK-SAME: lhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: lhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_precision_type = #vhlo.type_v1, ++// CHECK-SAME: num_primitive_operations = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++// CHECK-SAME: rhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: rhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_precision_type = #vhlo.type_v1 ++// CHECK-SAME: }> : (!vhlo.tensor_v1<8x8x16x!vhlo.f32_v1>, !vhlo.tensor_v1<8x16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<8x8x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) { ++ dot_dimension_numbers = #stablehlo.dot< ++ lhs_batching_dimensions = [0], ++ lhs_contracting_dimensions = [2], ++ rhs_batching_dimensions = [0], ++ rhs_contracting_dimensions = [1] ++ >, ++ algorithm = #stablehlo.dot_algorithm< ++ lhs_precision_type = tf32, ++ rhs_precision_type = tf32, ++ accumulation_type = f32, ++ lhs_component_count = 1, ++ rhs_component_count = 1, ++ num_primitive_operations = 1, ++ allow_imprecise_accumulation = false ++ > ++ } : (tensor<8x8x16xf32>, tensor<8x16x8xf32>) -> tensor<8x8x8xf32> ++ func.return %0 : tensor<8x8x8xf32> ++} ++ ++// CHECK-LABEL: "dot_general_algorithm_f8e4m3fn_x3" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @dot_general_algorithm_f8e4m3fn_x3(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { ++// CHECK: "vhlo.dot_general_v2"(%[[ARG0]], %[[ARG1]]) <{ ++// CHECK-SAME: accumulation_type = #vhlo.type_v1, ++// CHECK-SAME: allow_imprecise_accumulation = #vhlo.bool_v1, ++// CHECK-SAME: lhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: lhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_precision_type = #vhlo.type_v1, ++// CHECK-SAME: num_primitive_operations = #vhlo.integer_v1<3 : i64>, ++// CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++// CHECK-SAME: rhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: rhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_precision_type = #vhlo.type_v1 ++// CHECK-SAME: }> : (!vhlo.tensor_v1<8x8x16x!vhlo.bf16_v1>, !vhlo.tensor_v1<8x16x8x!vhlo.bf16_v1>) -> !vhlo.tensor_v1<8x8x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) { ++ dot_dimension_numbers = #stablehlo.dot< ++ lhs_batching_dimensions = [0], ++ lhs_contracting_dimensions = [2], ++ rhs_batching_dimensions = [0], ++ rhs_contracting_dimensions = [1] ++ >, ++ algorithm = #stablehlo.dot_algorithm< ++ lhs_precision_type = f8E4M3FN, ++ rhs_precision_type = f8E4M3FN, ++ accumulation_type = f32, ++ lhs_component_count = 1, ++ rhs_component_count = 1, ++ num_primitive_operations = 3, ++ allow_imprecise_accumulation = false ++ > ++ } : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> ++ func.return %0 : tensor<8x8x8xf32> ++} ++ ++// CHECK-LABEL: "dot_general_algorithm_f8e4m3fn_x4" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @dot_general_algorithm_f8e4m3fn_x4(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { ++// CHECK: "vhlo.dot_general_v2"(%[[ARG0]], %[[ARG1]]) <{ ++// CHECK-SAME: accumulation_type = #vhlo.type_v1, ++// CHECK-SAME: allow_imprecise_accumulation = #vhlo.bool_v1, ++// CHECK-SAME: lhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: lhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_precision_type = #vhlo.type_v1, ++// CHECK-SAME: num_primitive_operations = #vhlo.integer_v1<4 : i64>, ++// CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++// CHECK-SAME: rhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: rhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_precision_type = #vhlo.type_v1 ++// CHECK-SAME: }> : (!vhlo.tensor_v1<8x8x16x!vhlo.bf16_v1>, !vhlo.tensor_v1<8x16x8x!vhlo.bf16_v1>) -> !vhlo.tensor_v1<8x8x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) { ++ dot_dimension_numbers = #stablehlo.dot< ++ lhs_batching_dimensions = [0], ++ lhs_contracting_dimensions = [2], ++ rhs_batching_dimensions = [0], ++ rhs_contracting_dimensions = [1] ++ >, ++ algorithm = #stablehlo.dot_algorithm< ++ lhs_precision_type = f8E4M3FN, ++ rhs_precision_type = f8E4M3FN, ++ accumulation_type = f32, ++ lhs_component_count = 1, ++ rhs_component_count = 1, ++ num_primitive_operations = 4, ++ allow_imprecise_accumulation = false ++ > ++ } : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> ++ func.return %0 : tensor<8x8x8xf32> ++} ++ ++// CHECK-LABEL: "default_dynamic_broadcast_in_dim" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @default_dynamic_broadcast_in_dim(%arg0: tensor, %arg1: tensor<2xindex>) -> tensor { ++ // CHECK: "vhlo.dynamic_broadcast_in_dim_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: broadcast_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: known_expanding_dimensions = #vhlo.tensor_v1 : tensor<0xi64>>, ++ // CHECK-SAME: known_nonexpanding_dimensions = #vhlo.tensor_v1 : tensor<0xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1, !vhlo.tensor_v1<2x!vhlo.index_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.dynamic_broadcast_in_dim"(%arg0, %arg1) { ++ broadcast_dimensions = array ++ } : (tensor, tensor<2xindex>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "default_dynamic_conv" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @default_dynamic_conv(%arg0: tensor<1x8x8x207xf32>, %arg1: tensor<3x3x207x16xf32>, %arg2: tensor<2x2xi64>) -> tensor<1x?x?x16xf32> { ++ // CHECK: "vhlo.dynamic_conv_v2"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: batch_group_count = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: feature_group_count = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: input_batch_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: input_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: input_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: kernel_input_feature_dimension = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: kernel_output_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: kernel_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: lhs_dilation = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: output_batch_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: output_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: output_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++ // CHECK-SAME: rhs_dilation = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: window_reversal = #vhlo.tensor_v1 : tensor<2xi1>>, ++ // CHECK-SAME: window_strides = #vhlo.tensor_v1 : tensor<2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<1x8x8x207x!vhlo.f32_v1>, !vhlo.tensor_v1<3x3x207x16x!vhlo.f32_v1>, !vhlo.tensor_v1<2x2x!vhlo.i64_v1>) -> !vhlo.tensor_v1<1x?x?x16x!vhlo.f32_v1> ++ %0 = "stablehlo.dynamic_conv"(%arg0, %arg1, %arg2) { ++ dimension_numbers = #stablehlo.conv<[b, 0, 1, f]x[0, 1, i, o]->[b, 0, 1, f]>, ++ feature_group_count = 1 : i64, ++ batch_group_count = 1 : i64 ++ } : (tensor<1x8x8x207xf32>, tensor<3x3x207x16xf32>, tensor<2x2xi64>) -> tensor<1x?x?x16xf32> ++ func.return %0 : tensor<1x?x?x16xf32> ++} ++ ++// CHECK-LABEL: "default_dynamic_gather" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @default_dynamic_gather(%arg0 : tensor<2x4x9xf32>, %arg1 : tensor<1x5x2xi32>, %arg2 : tensor<3xi32>) -> tensor<1x5x8xf32> { ++ // CHECK: "vhlo.dynamic_gather_v2"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: collapsed_slice_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: index_vector_dim = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: indices_are_sorted = #vhlo.bool_v1, ++ // CHECK-SAME: offset_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: operand_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>>, ++ // CHECK-SAME: start_index_map = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: start_indices_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<2x4x9x!vhlo.f32_v1>, !vhlo.tensor_v1<1x5x2x!vhlo.i32_v1>, !vhlo.tensor_v1<3x!vhlo.i32_v1>) -> !vhlo.tensor_v1<1x5x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dynamic_gather"(%arg0, %arg1, %arg2) { ++ dimension_numbers = #stablehlo.gather< ++ offset_dims = [2], ++ collapsed_slice_dims = [0, 1], ++ start_index_map = [0, 1], ++ index_vector_dim = 2 ++ > ++ } : (tensor<2x4x9xf32>, tensor<1x5x2xi32>, tensor<3xi32>) -> tensor<1x5x8xf32> ++ func.return %0 : tensor<1x5x8xf32> ++} ++ ++func.func @default_func(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.func_v1"() <{ ++ // CHECK-SAME: arg_attrs = #vhlo.array_v1<[]>, ++ // CHECK-SAME: function_type = #vhlo.type_v1) -> !vhlo.tensor_v1>>, ++ // CHECK-SAME: res_attrs = #vhlo.array_v1<[]>, ++ // CHECK-SAME: sym_name = #vhlo.string_v1<"default_func">, ++ // CHECK-SAME: sym_visibility = #vhlo.string_v1<""> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG0:.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: "vhlo.return_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : () -> () ++ func.return %arg0 : tensor ++} ++ ++// CHECK-LABEL: "default_gather" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @default_gather(%arg0 : tensor<2x4x9xf32>, %arg1 : tensor<1x5x2xi32>) -> tensor<1x5x1xf32> { ++ // CHECK: "vhlo.gather_v2"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: collapsed_slice_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: index_vector_dim = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: indices_are_sorted = #vhlo.bool_v1, ++ // CHECK-SAME: offset_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: operand_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>>, ++ // CHECK-SAME: slice_sizes = #vhlo.tensor_v1 : tensor<3xi64>>, ++ // CHECK-SAME: start_index_map = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: start_indices_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<2x4x9x!vhlo.f32_v1>, !vhlo.tensor_v1<1x5x2x!vhlo.i32_v1>) -> !vhlo.tensor_v1<1x5x1x!vhlo.f32_v1> ++ %0 = "stablehlo.gather"(%arg0, %arg1) { ++ dimension_numbers = #stablehlo.gather< ++ offset_dims = [2], ++ collapsed_slice_dims = [0, 1], ++ start_index_map = [0, 1], ++ index_vector_dim = 2 ++ >, ++ slice_sizes = array ++ } : (tensor<2x4x9xf32>, tensor<1x5x2xi32>) -> tensor<1x5x1xf32> ++ func.return %0 : tensor<1x5x1xf32> ++} ++ ++// CHECK-LABEL: "default_infeed" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_infeed(%arg0: !stablehlo.token) -> (tensor, !stablehlo.token) { ++ // CHECK: "vhlo.infeed_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: infeed_config = #vhlo.string_v1<"">, ++ // CHECK-SAME{LITERAL}: layout = #vhlo.array_v1<[]> ++ // CHECK-SAME: }> : (!vhlo.token_v1) -> (!vhlo.tensor_v1, !vhlo.token_v1) ++ %0:2 = "stablehlo.infeed"(%arg0) : (!stablehlo.token) -> (tensor, !stablehlo.token) ++ func.return %0#0, %0#1 : tensor, !stablehlo.token ++} ++ ++// CHECK-LABEL: "default_outfeed" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @default_outfeed(%arg0: tensor, %arg1: !stablehlo.token) -> !stablehlo.token { ++ // CHECK: "vhlo.outfeed_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: outfeed_config = #vhlo.string_v1<""> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1, !vhlo.token_v1) -> !vhlo.token_v1 ++ %0 = "stablehlo.outfeed"(%arg0, %arg1) : (tensor, !stablehlo.token) -> !stablehlo.token ++ func.return %0 : !stablehlo.token ++} ++ ++// CHECK-LABEL: "op_recv_with_source_target_pairs" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_recv_with_source_target_pairs(%arg0: !stablehlo.token) -> (tensor, !stablehlo.token) { ++ // CHECK: "vhlo.recv_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: channel_type = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: is_host_transfer = #vhlo.bool_v1, ++ // CHECK-SAME{LITERAL}: source_target_pairs = #vhlo.tensor_v1 : tensor<2x2xi64>> ++ // CHECK-SAME{LITERAL}: }> : (!vhlo.token_v1) -> (!vhlo.tensor_v1, !vhlo.token_v1) ++ %0:2 = "stablehlo.recv"(%arg0) { ++ channel_handle = #stablehlo.channel_handle, ++ source_target_pairs = dense<[[0, 1], [1, 2]]> : tensor<2x2xi64> ++ } : (!stablehlo.token) -> (tensor, !stablehlo.token) ++ func.return %0#0, %0#1 : tensor, !stablehlo.token ++} ++ ++// CHECK-LABEL: "default_send" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @default_send(%arg0: tensor, %arg1: !stablehlo.token) -> !stablehlo.token { ++ // CHECK: "vhlo.send_v2"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: channel_type = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: is_host_transfer = #vhlo.bool_v1, ++ // CHECK-SAME{LITERAL}: source_target_pairs = #vhlo.tensor_v1 : tensor<2x2xi64>> ++ // CHECK-SAME{LITERAL}: }> : (!vhlo.tensor_v1, !vhlo.token_v1) -> !vhlo.token_v1 ++ %0 = "stablehlo.send"(%arg0, %arg1) { ++ channel_handle = #stablehlo.channel_handle, ++ source_target_pairs = dense<[[0, 1], [1, 2]]> : tensor<2x2xi64> ++ } : (tensor, !stablehlo.token) -> !stablehlo.token ++ func.return %0 : !stablehlo.token ++} ++ ++// CHECK-LABEL: "default_reduce_scatter" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_reduce_scatter(%arg0: tensor<16xf32>) -> tensor<16xf32> { ++ // CHECK: "vhlo.reduce_scatter_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<2x1xi64>>, ++ // CHECK-SAME: scatter_dimension = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.add_v1"(%[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.reduce_scatter"(%arg0) ({ ++ ^bb0(%arg1: tensor, %arg2: tensor): ++ %1 = "stablehlo.add"(%arg1, %arg2) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ scatter_dimension = 0 : i64, ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64> ++ } : (tensor<16xf32>) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "default_reduce_window" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @default_reduce_window(%arg0: tensor<2x17x31x7xf32>, %arg1: tensor) -> tensor<2x16x30x7xf32> { ++ // CHECK: "vhlo.reduce_window_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: base_dilations = #vhlo.tensor_v1 : tensor<4xi64>>, ++ // CHECK-SAME{LITERAL}: padding = #vhlo.tensor_v1 : tensor<4x2xi64>>, ++ // CHECK-SAME: window_dilations = #vhlo.tensor_v1 : tensor<4xi64>>, ++ // CHECK-SAME: window_dimensions = #vhlo.tensor_v1 : tensor<4xi64>>, ++ // CHECK-SAME: window_strides = #vhlo.tensor_v1 : tensor<4xi64>> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG2:arg.*]]: !vhlo.tensor_v1, %[[ARG3:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.maximum_v1"(%[[ARG2]], %[[ARG3]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<2x17x31x7x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1<2x16x30x7x!vhlo.f32_v1> ++ %0 = "stablehlo.reduce_window"(%arg0, %arg1) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor): ++ %1 = "stablehlo.maximum"(%arg2, %arg3) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ window_dimensions = array ++ } : (tensor<2x17x31x7xf32>, tensor) -> tensor<2x16x30x7xf32> ++ func.return %0 : tensor<2x16x30x7xf32> ++} ++ ++// CHECK-LABEL: "default_scatter" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @default_scatter(%arg0: tensor<200x100x300xf32>, %arg1: tensor<10x2xi32>, %arg2: tensor<10x300xf32>) -> tensor<200x100x300xf32> { ++ // CHECK: "vhlo.scatter_v2"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: index_vector_dim = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: indices_are_sorted = #vhlo.bool_v1, ++ // CHECK-SAME: input_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>>, ++ // CHECK-SAME: inserted_window_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: scatter_dims_to_operand_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: scatter_indices_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>>, ++ // CHECK-SAME: unique_indices = #vhlo.bool_v1, ++ // CHECK-SAME: update_window_dims = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG3:arg.*]]: !vhlo.tensor_v1, %[[ARG4:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.add_v1"(%[[ARG3]], %[[ARG4]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<200x100x300x!vhlo.f32_v1>, !vhlo.tensor_v1<10x2x!vhlo.i32_v1>, !vhlo.tensor_v1<10x300x!vhlo.f32_v1>) -> !vhlo.tensor_v1<200x100x300x!vhlo.f32_v1> ++ %0 = "stablehlo.scatter"(%arg0, %arg1, %arg2) ({ ++ ^bb0(%arg3: tensor, %arg4: tensor): ++ %1 = "stablehlo.add"(%arg3, %arg4) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ scatter_dimension_numbers = #stablehlo.scatter< ++ update_window_dims = [1], ++ inserted_window_dims = [0, 1], ++ scatter_dims_to_operand_dims = [0, 1], ++ index_vector_dim = 1 ++ > ++ } : (tensor<200x100x300xf32>, tensor<10x2xi32>, tensor<10x300xf32>) -> tensor<200x100x300xf32> ++ func.return %0 : tensor<200x100x300xf32> ++} ++ ++// CHECK-LABEL: "default_select_and_scatter" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @default_select_and_scatter(%arg0: tensor<10x24x24x64xf32>, %arg1: tensor<10x23x23x64xf32>, %arg2: tensor) -> tensor<10x24x24x64xf32> { ++ // CHECK: "vhlo.select_and_scatter_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: padding = #vhlo.tensor_v1 : tensor<4x2xi64>>, ++ // CHECK-SAME: window_dimensions = #vhlo.tensor_v1 : tensor<4xi64>>, ++ // CHECK-SAME: window_strides = #vhlo.tensor_v1 : tensor<4xi64>> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG31:arg.*]]: !vhlo.tensor_v1, %[[ARG41:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL11:.*]] = "vhlo.compare_v1"(%[[ARG31]], %[[ARG41]]) <{compare_type = #vhlo, comparison_direction = #vhlo}> ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL11]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }, { ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG32:arg.*]]: !vhlo.tensor_v1, %[[ARG42:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL12:.*]] = "vhlo.add_v1"(%[[ARG32]], %[[ARG42]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL12]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<10x24x24x64x!vhlo.f32_v1>, !vhlo.tensor_v1<10x23x23x64x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1<10x24x24x64x!vhlo.f32_v1> ++ %0 = "stablehlo.select_and_scatter"(%arg0, %arg1, %arg2) ({ ++ ^bb0(%arg3: tensor, %arg4: tensor): ++ %1 = "stablehlo.compare"(%arg3, %arg4) {compare_type = #stablehlo, comparison_direction = #stablehlo} : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }, { ++ ^bb0(%arg3: tensor, %arg4: tensor): ++ %1 = "stablehlo.add"(%arg3, %arg4) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ window_dimensions = array ++ } : (tensor<10x24x24x64xf32>, tensor<10x23x23x64xf32>, tensor) -> tensor<10x24x24x64xf32> ++ func.return %0 : tensor<10x24x24x64xf32> ++} ++ ++// CHECK-LABEL: "default_sort" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @default_sort(%arg0: tensor<16xf32>) -> tensor<16xf32> { ++ // CHECK: "vhlo.sort_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: dimension = #vhlo.integer_v1<-1 : i64> ++ // CHECK-SAME: is_stable = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.compare_v1"(%[[ARG1]], %[[ARG2]]) <{compare_type = #vhlo, comparison_direction = #vhlo}> ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.sort"(%arg0) ({ ++ ^bb0(%arg1: tensor, %arg2: tensor): ++ %1 = "stablehlo.compare"(%arg1, %arg2) {compare_type = #stablehlo, comparison_direction = #stablehlo} : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) : (tensor<16xf32>) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// ============ OPS ============ ++ ++// CHECK-LABEL: "op_abs" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_abs(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.abs_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.abs"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_add" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_add(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_after_all" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_after_all(%arg0: !stablehlo.token) -> !stablehlo.token { ++ // CHECK: "vhlo.after_all_v1"(%[[ARG0]]) : (!vhlo.token_v1) -> !vhlo.token_v1 ++ %0 = "stablehlo.after_all"(%arg0) : (!stablehlo.token) -> !stablehlo.token ++ func.return %0 : !stablehlo.token ++} ++ ++// CHECK-LABEL: "op_all_gather" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_all_gather(%arg0: tensor<16x8xf32>) -> tensor<16x16xf32> { ++ // CHECK: "vhlo.all_gather_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: all_gather_dim = #vhlo.integer_v1<1 : i64> ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<2x1xi64>>, ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x16x!vhlo.f32_v1> ++ %0 = "stablehlo.all_gather"(%arg0) { ++ all_gather_dim = 1 : i64, ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64>, ++ channel_handle = #stablehlo.channel_handle, ++ use_global_device_ids ++ } : (tensor<16x8xf32>) -> tensor<16x16xf32> ++ func.return %0 : tensor<16x16xf32> ++} ++ ++// CHECK-LABEL: "op_all_reduce" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_all_reduce(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.all_reduce_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<2x1xi64>>, ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.add_v1"(%[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.all_reduce"(%arg0) ({ ++ ^bb0(%arg1: tensor, %arg2: tensor): ++ %1 = "stablehlo.add"(%arg1, %arg2) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64>, ++ channel_handle = #stablehlo.channel_handle, ++ use_global_device_ids ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_all_reduce_with_promotable_types" ++func.func @op_all_reduce_with_promotable_types(%operand: tensor) -> tensor { ++ // CHECK: "vhlo.all_reduce_v2"(%[[ARG0:.*]]) ++ // CHECK: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK: "vhlo.return_v1"(%[[VAL1:.*]]) : (!vhlo.tensor_v1) -> () ++ // CHECK: }) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %result = "stablehlo.all_reduce"(%operand) ({ ++ ^bb0(%arg0: tensor, %arg1: tensor): ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%0) : (tensor) -> () ++ }) { ++ replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>, ++ channel_handle = #stablehlo.channel_handle, ++ use_global_device_ids ++ } : (tensor) -> tensor ++ ++ func.return %result : tensor ++} ++ ++// CHECK-LABEL: "default_all_reduce_variadic" ++func.func @default_all_reduce_variadic(%arg0: tensor, %arg1: tensor) -> (tensor, tensor) { ++ %0:2 = "stablehlo.all_reduce"(%arg0, %arg1) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor): ++ %1 = "stablehlo.add"(%arg2, %arg3) : (tensor, tensor) -> (tensor) ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64> ++ } : (tensor, tensor) -> (tensor, tensor) ++ func.return %0#0, %0#1 : tensor, tensor ++} ++ ++// CHECK-LABEL: "op_all_to_all" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_all_to_all(%arg0: tensor<4x16xf32>) -> tensor<16x4xf32> { ++ // CHECK: "vhlo.all_to_all_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: concat_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<1x4xi64>>, ++ // CHECK-SAME: split_count = #vhlo.integer_v1<4 : i64> ++ // CHECK-SAME: split_dimension = #vhlo.integer_v1<1 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<4x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x4x!vhlo.f32_v1> ++ %0 = "stablehlo.all_to_all"(%arg0) { ++ split_dimension = 1 : i64, ++ concat_dimension = 0 : i64, ++ split_count = 4 : i64, ++ replica_groups = dense<[[0, 1, 2, 3]]> : tensor<1x4xi64>, ++ channel_handle = #stablehlo.channel_handle ++ } : (tensor<4x16xf32>) -> tensor<16x4xf32> ++ func.return %0 : tensor<16x4xf32> ++} ++ ++// CHECK-LABEL: "op_and" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_and(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.and_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.and"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_atan2" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_atan2(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.atan2_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.atan2"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_batch_norm_grad" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}, %[[ARG3:.*]]: {{.*}}, %[[ARG4:.*]]: {{.*}}) ++func.func @op_batch_norm_grad(%arg0: tensor<16x16x16x16xf32>, %arg1: tensor<16xf32>, %arg2: tensor<16xf32>, %arg3: tensor<16xf32>, %arg4: tensor<16x16x16x16xf32>) -> (tensor<16x16x16x16xf32>, tensor<16xf32>, tensor<16xf32>) { ++ // CHECK: "vhlo.batch_norm_grad_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]], %[[ARG3]], %[[ARG4]]) <{ ++ // CHECK-SAME: epsilon = #vhlo.float_v1<1.000000e-03 : !vhlo.f32_v1>, ++ // CHECK-SAME: feature_index = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x16x16x16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x16x16x16x!vhlo.f32_v1>) -> (!vhlo.tensor_v1<16x16x16x16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>) ++ %0:3 = "stablehlo.batch_norm_grad"(%arg0, %arg1, %arg2, %arg3, %arg4) { ++ epsilon = 0.001 : f32, ++ feature_index = 0 : i64 ++ } : (tensor<16x16x16x16xf32>, tensor<16xf32>, tensor<16xf32>, tensor<16xf32>, tensor<16x16x16x16xf32>) -> (tensor<16x16x16x16xf32>, tensor<16xf32>, tensor<16xf32>) ++ func.return %0#0, %0#1, %0#2 : tensor<16x16x16x16xf32>, tensor<16xf32>, tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_batch_norm_inference" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}, %[[ARG3:.*]]: {{.*}}, %[[ARG4:.*]]: {{.*}}) ++func.func @op_batch_norm_inference(%arg0: tensor<16x16x16x16xf32>, %arg1: tensor<16xf32>, %arg2: tensor<16xf32>, %arg3: tensor<16xf32>, %arg4: tensor<16xf32>) -> tensor<16x16x16x16xf32> { ++ // CHECK: "vhlo.batch_norm_inference_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]], %[[ARG3]], %[[ARG4]]) <{ ++ // CHECK-SAME: epsilon = #vhlo.float_v1<1.000000e-03 : !vhlo.f32_v1>, ++ // CHECK-SAME: feature_index = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x16x16x16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x16x16x16x!vhlo.f32_v1> ++ %0 = "stablehlo.batch_norm_inference"(%arg0, %arg1, %arg2, %arg3, %arg4) { ++ epsilon = 0.001 : f32, ++ feature_index = 0 : i64 ++ } : (tensor<16x16x16x16xf32>, tensor<16xf32>, tensor<16xf32>, tensor<16xf32>, tensor<16xf32>) -> tensor<16x16x16x16xf32> ++ func.return %0 : tensor<16x16x16x16xf32> ++} ++ ++// CHECK-LABEL: "op_batch_norm_training" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_batch_norm_training(%arg0: tensor<16x16x16x16xf32>, %arg1: tensor<16xf32>, %arg2: tensor<16xf32>) -> (tensor<16x16x16x16xf32>, tensor<16xf32>, tensor<16xf32>) { ++ // CHECK: "vhlo.batch_norm_training_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: epsilon = #vhlo.float_v1<1.000000e-03 : !vhlo.f32_v1>, ++ // CHECK-SAME: feature_index = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x16x16x16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>) -> (!vhlo.tensor_v1<16x16x16x16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x!vhlo.f32_v1>) ++ %0:3 = "stablehlo.batch_norm_training"(%arg0, %arg1, %arg2) { ++ epsilon = 0.001 : f32, ++ feature_index = 0 : i64 ++ } : (tensor<16x16x16x16xf32>, tensor<16xf32>, tensor<16xf32>) -> (tensor<16x16x16x16xf32>, tensor<16xf32>, tensor<16xf32>) ++ func.return %0#0, %0#1, %0#2 : tensor<16x16x16x16xf32>, tensor<16xf32>, tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_bitcast_convert" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_bitcast_convert(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.bitcast_convert_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.bitcast_convert"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_broadcast_in_dim" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_broadcast_in_dim(%arg0: tensor<16xf32>) -> tensor<16x16xf32> { ++ // CHECK: "vhlo.broadcast_in_dim_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: broadcast_dimensions = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x16x!vhlo.f32_v1> ++ %0 = "stablehlo.broadcast_in_dim"(%arg0) { ++ broadcast_dimensions = array ++ } : (tensor<16xf32>) -> tensor<16x16xf32> ++ func.return %0 : tensor<16x16xf32> ++} ++ ++// CHECK-LABEL: "op_broadcast" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_broadcast(%arg0: tensor<16xf32>) -> tensor<16x16xf32> { ++ // CHECK: "vhlo.broadcast_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: broadcast_sizes = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x16x!vhlo.f32_v1> ++ %0 = "stablehlo.broadcast"(%arg0) { ++ broadcast_sizes = array ++ } : (tensor<16xf32>) -> tensor<16x16xf32> ++ func.return %0 : tensor<16x16xf32> ++} ++ ++// CHECK-LABEL: "op_case" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_case(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.case_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: "vhlo.return_v1"(%[[ARG1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.case"(%arg0) ({ ++ "stablehlo.return"(%arg1) : (tensor) -> () ++ }) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_cbrt" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_cbrt(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.cbrt_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.cbrt"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_ceil" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_ceil(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.ceil_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.ceil"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_cholesky" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_cholesky(%arg0: tensor<1x16x16xf32>) -> tensor<1x16x16xf32> { ++ // CHECK: "vhlo.cholesky_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: lower = #vhlo.bool_v1 ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<1x16x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<1x16x16x!vhlo.f32_v1> ++ %0 = "stablehlo.cholesky"(%arg0) { ++ lower = true ++ } : (tensor<1x16x16xf32>) -> tensor<1x16x16xf32> ++ func.return %0 : tensor<1x16x16xf32> ++} ++ ++// CHECK-LABEL: "op_clamp" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_clamp(%arg0: tensor, %arg1: tensor, %arg2: tensor) -> tensor { ++ // CHECK: "vhlo.clamp_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.clamp"(%arg0, %arg1, %arg2) : (tensor, tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_count_leading_zeros" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_count_leading_zeros(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.count_leading_zeros_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.count_leading_zeros"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_collective_broadcast" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_collective_broadcast(%arg0: tensor<16x8xf32>, %arg1: tensor<1xi32>) -> tensor<16x8xf32> { ++ // CHECK: "vhlo.collective_broadcast_v2"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: has_dynamic_root = #vhlo.bool_v1, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<1x2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>, !vhlo.tensor_v1<1x!vhlo.i32_v1>) -> !vhlo.tensor_v1<16x8x!vhlo.f32_v1> ++ %0 = "stablehlo.collective_broadcast"(%arg0, %arg1) { ++ replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>, ++ channel_handle = #stablehlo.channel_handle, ++ has_dynamic_root ++ } : (tensor<16x8xf32>, tensor<1xi32>) -> tensor<16x8xf32> ++ func.return %0 : tensor<16x8xf32> ++} ++ ++// CHECK-LABEL: "op_collective_reduce" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_collective_reduce(%arg0: tensor, %arg1: tensor<1xi32>) -> tensor { ++ // CHECK: "vhlo.collective_reduce_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: has_dynamic_root = #vhlo.bool_v1, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<2x1xi64>>, ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG2:arg.*]]: !vhlo.tensor_v1, %[[ARG3:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.add_v1"(%[[ARG2]], %[[ARG3]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1, !vhlo.tensor_v1<1x!vhlo.i32_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.collective_reduce"(%arg0, %arg1) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor): ++ %1 = "stablehlo.add"(%arg2, %arg3) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64>, ++ channel_handle = #stablehlo.channel_handle, ++ use_global_device_ids, ++ has_dynamic_root ++ } : (tensor, tensor<1xi32>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_collective_permute" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_collective_permute(%arg0: tensor<16x8xf32>) -> tensor<16x8xf32> { ++ // CHECK: "vhlo.collective_permute_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME{LITERAL}: source_target_pairs = #vhlo.tensor_v1 : tensor<3x2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x8x!vhlo.f32_v1> ++ %0 = "stablehlo.collective_permute"(%arg0) { ++ source_target_pairs = dense<[[0, 1], [1, 2], [2, 3]]> : tensor<3x2xi64>, ++ channel_handle = #stablehlo.channel_handle ++ } : (tensor<16x8xf32>) -> tensor<16x8xf32> ++ func.return %0 : tensor<16x8xf32> ++} ++ ++// CHECK-LABEL: "op_compare" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_compare(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.compare_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: compare_type = #vhlo, ++ // CHECK-SAME: comparison_direction = #vhlo ++ // CHECK-SAME: }> : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.compare"(%arg0, %arg1) { ++ comparison_direction = #stablehlo, ++ compare_type = #stablehlo ++ } : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_complex" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_complex(%arg0: tensor, %arg1: tensor) -> tensor> { ++ // CHECK: "vhlo.complex_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1> ++ %0 = "stablehlo.complex"(%arg0, %arg1) : (tensor, tensor) -> tensor> ++ func.return %0 : tensor> ++} ++ ++// CHECK-LABEL: "op_composite" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_composite(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.composite_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: composite_attributes = #vhlo.dict_v1<{#vhlo.string_v1<"my_int"> = #vhlo.integer_v1<1 : i64>, #vhlo.string_v1<"my_string"> = #vhlo.string_v1<"foo">}> ++ // CHECK-SAME: decomposition = #vhlo.string_v1<"composite_target"> ++ // CHECK-SAME: name = #vhlo.string_v1<"stablehlo.composite_target"> ++ // CHECK-SAME: version = #vhlo.integer_v1<1 : i32> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.composite"(%arg0) { ++ name = "stablehlo.composite_target", ++ decomposition = @composite_target, ++ version = 1 : i32, ++ composite_attributes = { ++ my_int = 1 : i64, ++ my_string = "foo" ++ } ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "composite_regions" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @composite_regions(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.composite_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: composite_attributes = #vhlo.dict_v1<{}> ++ // CHECK-SAME: decomposition = #vhlo.string_v1<"composite_target"> ++ // CHECK-SAME: name = #vhlo.string_v1<"stablehlo.composite_target"> ++ // CHECK-SAME: version = #vhlo.integer_v1<1 : i32> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^bb0(%[[ARG1:.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: "vhlo.return_v1"(%[[ARG1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.composite"(%arg0) ({ ++ ^bb0(%arg1: tensor): ++ "stablehlo.return"(%arg1) : (tensor) -> () ++ }) { ++ name = "stablehlo.composite_target", ++ decomposition = @composite_target, ++ version = 1 : i32 ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_concatenate" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_concatenate(%arg0: tensor<8xf32>, %arg1: tensor<8xf32>) -> tensor<16xf32> { ++ // CHECK: "vhlo.concatenate_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: dimension = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<8x!vhlo.f32_v1>, !vhlo.tensor_v1<8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.concatenate"(%arg0, %arg1) { ++ dimension = 0 : i64 ++ } : (tensor<8xf32>, tensor<8xf32>) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_constant" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_constant(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.constant_v1"() <{ ++ // CHECK-SAME: value = #vhlo.tensor_v1 : tensor> ++ // CHECK-SAME: }> : () -> !vhlo.tensor_v1 ++ %0 = "stablehlo.constant"() { ++ value = dense<0.0> : tensor ++ } : () -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_convert" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_convert(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.convert_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.convert"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_convolution" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_convolution(%arg0: tensor<1x8x8x207xf32>, %arg1: tensor<3x3x207x16xf32>) -> tensor<1x7x7x16xf32> { ++ // CHECK: "vhlo.convolution_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: batch_group_count = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: feature_group_count = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: input_batch_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: input_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: input_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: kernel_input_feature_dimension = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: kernel_output_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: kernel_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: lhs_dilation = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: output_batch_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: output_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: output_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: padding = #vhlo.tensor_v1 : tensor<2x2xi64>>, ++ // CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++ // CHECK-SAME: rhs_dilation = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: window_reversal = #vhlo.tensor_v1 : tensor<2xi1>>, ++ // CHECK-SAME: window_strides = #vhlo.tensor_v1 : tensor<2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<1x8x8x207x!vhlo.f32_v1>, !vhlo.tensor_v1<3x3x207x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<1x7x7x16x!vhlo.f32_v1> ++ %0 = "stablehlo.convolution"(%arg0, %arg1) { ++ window_strides = array, ++ padding = dense<1> : tensor<2x2xi64>, ++ lhs_dilation = array, ++ rhs_dilation = array, ++ window_reversal = array, ++ dimension_numbers = #stablehlo.conv<[b, 0, 1, f]x[0, 1, i, o]->[b, 0, 1, f]>, ++ feature_group_count = 1 : i64, ++ batch_group_count = 1 : i64, ++ precision_config = [#stablehlo, #stablehlo] ++ } : (tensor<1x8x8x207xf32>, tensor<3x3x207x16xf32>) -> tensor<1x7x7x16xf32> ++ func.return %0 : tensor<1x7x7x16xf32> ++} ++ ++// CHECK-LABEL: "op_cosine" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_cosine(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.cosine_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.cosine"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_create_token" ++func.func @op_create_token() -> !stablehlo.token { ++ // CHECK: "vhlo.create_token_v1"() : () -> !vhlo.token_v1 ++ %0 = "stablehlo.create_token"() : () -> !stablehlo.token ++ func.return %0 : !stablehlo.token ++} ++ ++// CHECK-LABEL: "op_cross_replica_sum" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_cross_replica_sum(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.cross-replica-sum_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<2x1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.cross-replica-sum"(%arg0) { ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64> ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_custom_call" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_custom_call(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.custom_call_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: api_version = #vhlo, ++ // CHECK-SAME: backend_config = #vhlo.string_v1<"\08\03\1A\02">, ++ // CHECK-SAME: call_target_name = #vhlo.string_v1<"foo">, ++ // CHECK-SAME: called_computations = #vhlo.array_v1<[#vhlo.string_v1<"foo">]>, ++ // CHECK-SAME: has_side_effect = #vhlo.bool_v1, ++ // CHECK-SAME: operand_layouts = #vhlo.array_v1<[#vhlo.tensor_v1 : tensor<0xindex>>]>, ++ // CHECK-SAME: output_operand_aliases = #vhlo.array_v1<[ ++ // CHECK-SAME: #vhlo.output_operand_alias_v1< ++ // CHECK-SAME: outputTupleIndices = [], ++ // CHECK-SAME: operandIndex = 0, ++ // CHECK-SAME: operandTupleIndices = []>]>, ++ // CHECK-SAME: result_layouts = #vhlo.array_v1<[#vhlo.tensor_v1 : tensor<0xindex>>]>, ++ // CHECK-SAME: result_tilings = #vhlo.array_v1<[]> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "foo", ++ has_side_effect = true, ++ backend_config = "\08\03\1A\02", ++ api_version = 2 : i32, ++ called_computations = [@foo], ++ operand_layouts = [dense<> : tensor<0xindex>], ++ output_operand_aliases = [ ++ #stablehlo.output_operand_alias], ++ result_layouts = [dense<> : tensor<0xindex>] ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_custom_call_empty_result_layout" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func public @op_custom_call_empty_result_layout(%arg0: tensor) -> tensor { ++ // %0 = "vhlo.custom_call_v2"(%arg0) <{>}> : (!vhlo.tensor_v1) -> !vhlo.tuple_v1<> ++ // CHECK: "vhlo.custom_call_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: api_version = #vhlo, ++ // CHECK-SAME: backend_config = #vhlo.string_v1<"">, ++ // CHECK-SAME: call_target_name = #vhlo.string_v1<"empty_output">, ++ // CHECK-SAME: called_computations = #vhlo.array_v1<[]>, ++ // CHECK-SAME: has_side_effect = #vhlo.bool_v1, ++ // CHECK-SAME: operand_layouts = #vhlo.array_v1<[#vhlo.tensor_v1 : tensor<0xindex>>]>, ++ // CHECK-SAME: output_operand_aliases = #vhlo.array_v1<[]>, ++ // CHECK-SAME: result_layouts = #vhlo.array_v1<[]>, ++ // CHECK-SAME: result_tilings = #vhlo.array_v1<[]> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1) -> !vhlo.tuple_v1<> ++ %0 = "stablehlo.custom_call"(%arg0) <{ ++ api_version = 2 : i32, ++ call_target_name = "empty_output", ++ has_side_effect = true, ++ operand_layouts = [dense<> : tensor<0xindex>], ++ result_layouts = [] ++ }> : (tensor) -> tuple<> ++ return %arg0 : tensor ++} ++ ++// CHECK-LABEL: "op_custom_call_with_result_tilings" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_custom_call_with_result_tilings(%arg0: tensor<64x256xf32>) -> tensor<64x256xf32> { ++ // CHECK: "vhlo.custom_call_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: api_version = #vhlo, ++ // CHECK-SAME: backend_config = #vhlo.string_v1<"">, ++ // CHECK-SAME: call_target_name = #vhlo.string_v1<"foo">, ++ // CHECK-SAME: called_computations = #vhlo.array_v1<[]>, ++ // CHECK-SAME: has_side_effect = #vhlo.bool_v1, ++ // CHECK-SAME: operand_layouts = #vhlo.array_v1<[#vhlo.tensor_v1 : tensor<2xindex>>]>, ++ // CHECK-SAME: output_operand_aliases = #vhlo.array_v1<[]>, ++ // CHECK-SAME: result_layouts = #vhlo.array_v1<[#vhlo.tensor_v1 : tensor<2xindex>>]>, ++ // CHECK-SAME: result_tilings = #vhlo.array_v1<[#vhlo.array_v1<[#vhlo.tensor_v1 : tensor<2xindex>>]>]> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<64x256x!vhlo.f32_v1>) -> !vhlo.tensor_v1<64x256x!vhlo.f32_v1> ++ %0 = "stablehlo.custom_call"(%arg0) { ++ api_version = 2 : i32, ++ call_target_name = "foo", ++ operand_layouts = [dense<[1, 0]> : tensor<2xindex>], ++ result_layouts = [dense<[1, 0]> : tensor<2xindex>], ++ result_tilings = [[dense<[8, 128]> : tensor<2xindex>]] ++ } : (tensor<64x256xf32>) -> tensor<64x256xf32> ++ func.return %0 : tensor<64x256xf32> ++} ++ ++// CHECK-LABEL: "op_divide" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_divide(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.divide_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.divide"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_dot_general" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_dot_general(%arg0: tensor<8x8x16xf32>, %arg1: tensor<8x16x8xf32>) -> tensor<8x8x8xf32> { ++ // CHECK: "vhlo.dot_general_v2"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: accumulation_type = #vhlo.type_v1, ++ // CHECK-SAME: allow_imprecise_accumulation = #vhlo.type_v1, ++ // CHECK-SAME: lhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: lhs_component_count = #vhlo.type_v1, ++ // CHECK-SAME: lhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: lhs_precision_type = #vhlo.type_v1, ++ // CHECK-SAME: num_primitive_operations = #vhlo.type_v1, ++ // CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++ // CHECK-SAME: rhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: rhs_component_count = #vhlo.type_v1, ++ // CHECK-SAME: rhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: rhs_precision_type = #vhlo.type_v1 ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<8x8x16x!vhlo.f32_v1>, !vhlo.tensor_v1<8x16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<8x8x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) { ++ dot_dimension_numbers = #stablehlo.dot< ++ lhs_batching_dimensions = [0], ++ lhs_contracting_dimensions = [2], ++ rhs_batching_dimensions = [0], ++ rhs_contracting_dimensions = [1] ++ >, ++ precision_config = [#stablehlo, #stablehlo] ++ } : (tensor<8x8x16xf32>, tensor<8x16x8xf32>) -> tensor<8x8x8xf32> ++ func.return %0 : tensor<8x8x8xf32> ++} ++ ++// CHECK-LABEL: "op_dot" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_dot(%arg0: tensor<8x16xf32>, %arg1: tensor<16x8xf32>) -> tensor<8x8xf32> { ++ // CHECK: "vhlo.dot_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<8x16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<8x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dot"(%arg0, %arg1) { ++ precision_config = [#stablehlo, #stablehlo] ++ } : (tensor<8x16xf32>, tensor<16x8xf32>) -> tensor<8x8xf32> ++ func.return %0 : tensor<8x8xf32> ++} ++ ++// CHECK-LABEL: "op_dynamic_broadcast_in_dim" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_dynamic_broadcast_in_dim(%arg0: tensor, %arg1: tensor<2xindex>) -> tensor { ++ // CHECK: "vhlo.dynamic_broadcast_in_dim_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: broadcast_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: known_expanding_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: known_nonexpanding_dimensions = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1, !vhlo.tensor_v1<2x!vhlo.index_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.dynamic_broadcast_in_dim"(%arg0, %arg1) { ++ broadcast_dimensions = array, ++ known_expanding_dimensions = array, ++ known_nonexpanding_dimensions = array ++ } : (tensor, tensor<2xindex>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_dynamic_conv" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_dynamic_conv(%arg0: tensor<1x8x8x207xf32>, %arg1: tensor<3x3x207x16xf32>, %arg2: tensor<2x2xi64>) -> tensor<1x?x?x16xf32> { ++ // CHECK: "vhlo.dynamic_conv_v2"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: batch_group_count = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: feature_group_count = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: input_batch_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: input_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: input_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: kernel_input_feature_dimension = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: kernel_output_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: kernel_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: lhs_dilation = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: output_batch_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: output_feature_dimension = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: output_spatial_dimensions = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++ // CHECK-SAME: rhs_dilation = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: window_reversal = #vhlo.tensor_v1 : tensor<2xi1>>, ++ // CHECK-SAME: window_strides = #vhlo.tensor_v1 : tensor<2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<1x8x8x207x!vhlo.f32_v1>, !vhlo.tensor_v1<3x3x207x16x!vhlo.f32_v1>, !vhlo.tensor_v1<2x2x!vhlo.i64_v1>) -> !vhlo.tensor_v1<1x?x?x16x!vhlo.f32_v1> ++ %0 = "stablehlo.dynamic_conv"(%arg0, %arg1, %arg2) { ++ window_strides = array, ++ lhs_dilation = array, ++ rhs_dilation = array, ++ window_reversal = array, ++ dimension_numbers = #stablehlo.conv<[b, 0, 1, f]x[0, 1, i, o]->[b, 0, 1, f]>, ++ feature_group_count = 1 : i64, ++ batch_group_count = 1 : i64, ++ precision_config = [#stablehlo, #stablehlo] ++ } : (tensor<1x8x8x207xf32>, tensor<3x3x207x16xf32>, tensor<2x2xi64>) -> tensor<1x?x?x16xf32> ++ func.return %0 : tensor<1x?x?x16xf32> ++} ++ ++// CHECK-LABEL: "op_dynamic_gather" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_dynamic_gather(%arg0 : tensor<2x4x9xf32>, %arg1 : tensor<1x5x2xi32>, %arg2 : tensor<3xi32>) -> tensor<1x5x8xf32> { ++ // CHECK: "vhlo.dynamic_gather_v2"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: collapsed_slice_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: index_vector_dim = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: indices_are_sorted = #vhlo.bool_v1, ++ // CHECK-SAME: offset_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: operand_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>>, ++ // CHECK-SAME: start_index_map = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: start_indices_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<2x4x9x!vhlo.f32_v1>, !vhlo.tensor_v1<1x5x2x!vhlo.i32_v1>, !vhlo.tensor_v1<3x!vhlo.i32_v1>) -> !vhlo.tensor_v1<1x5x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dynamic_gather"(%arg0, %arg1, %arg2) { ++ dimension_numbers = #stablehlo.gather< ++ offset_dims = [2], ++ collapsed_slice_dims = [0, 1], ++ start_index_map = [0, 1], ++ index_vector_dim = 2 ++ >, ++ indices_are_sorted = true ++ } : (tensor<2x4x9xf32>, tensor<1x5x2xi32>, tensor<3xi32>) -> tensor<1x5x8xf32> ++ func.return %0 : tensor<1x5x8xf32> ++} ++ ++// CHECK-LABEL: "op_dynamic_gather_with_batching_dims" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_dynamic_gather_with_batching_dims(%arg0 : tensor<5x2x4x9xf32>, %arg1 : tensor<1x5x2xi32>, %arg2 : tensor<4xi32>) -> tensor<1x5x8xf32> { ++ // CHECK: "vhlo.dynamic_gather_v2"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: collapsed_slice_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: index_vector_dim = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: indices_are_sorted = #vhlo.bool_v1, ++ // CHECK-SAME: offset_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: operand_batching_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: start_index_map = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: start_indices_batching_dims = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<5x2x4x9x!vhlo.f32_v1>, !vhlo.tensor_v1<1x5x2x!vhlo.i32_v1>, !vhlo.tensor_v1<4x!vhlo.i32_v1>) -> !vhlo.tensor_v1<1x5x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dynamic_gather"(%arg0, %arg1, %arg2) { ++ dimension_numbers = #stablehlo.gather< ++ offset_dims = [2], ++ collapsed_slice_dims = [1, 2], ++ operand_batching_dims = [0], ++ start_indices_batching_dims = [1], ++ start_index_map = [1, 2], ++ index_vector_dim = 2 ++ >, ++ indices_are_sorted = true ++ } : (tensor<5x2x4x9xf32>, tensor<1x5x2xi32>, tensor<4xi32>) -> tensor<1x5x8xf32> ++ func.return %0 : tensor<1x5x8xf32> ++} ++ ++// CHECK-LABEL: "op_dynamic_iota" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_dynamic_iota(%arg0: tensor<1xindex>) -> tensor { ++ // CHECK: "vhlo.dynamic_iota_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: iota_dimension = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<1x!vhlo.index_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.dynamic_iota"(%arg0) { ++ iota_dimension = 0 : i64 ++ } : (tensor<1xindex>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_dynamic_pad" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}, %[[ARG3:.*]]: {{.*}}, %[[ARG4:.*]]: {{.*}}) ++func.func @op_dynamic_pad(%arg0: tensor, %arg1: tensor, %arg2: tensor<1xindex>, %arg3: tensor<1xindex>, %arg4: tensor<1xindex>) -> tensor { ++ // CHECK: "vhlo.dynamic_pad_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]], %[[ARG3]], %[[ARG4]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1, !vhlo.tensor_v1<1x!vhlo.index_v1>, !vhlo.tensor_v1<1x!vhlo.index_v1>, !vhlo.tensor_v1<1x!vhlo.index_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.dynamic_pad"(%arg0, %arg1, %arg2, %arg3, %arg4) : (tensor, tensor, tensor<1xindex>, tensor<1xindex>, tensor<1xindex>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_dynamic_reshape" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_dynamic_reshape(%arg0: tensor<16xf32>, %arg1: tensor<2xindex>) -> tensor { ++ // CHECK: "vhlo.dynamic_reshape_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<2x!vhlo.index_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.dynamic_reshape"(%arg0, %arg1) : (tensor<16xf32>, tensor<2xindex>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_dynamic_slice" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_dynamic_slice(%arg0: tensor<16xf32>, %arg1: tensor) -> tensor<4xf32> { ++ // CHECK: "vhlo.dynamic_slice_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: slice_sizes = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1<4x!vhlo.f32_v1> ++ %0 = "stablehlo.dynamic_slice"(%arg0, %arg1) { ++ slice_sizes = array ++ } : (tensor<16xf32>, tensor) -> tensor<4xf32> ++ func.return %0 : tensor<4xf32> ++} ++ ++// CHECK-LABEL: "op_dynamic_update_slice" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_dynamic_update_slice(%arg0: tensor<16xf32>, %arg1: tensor<4xf32>, %arg2: tensor) -> tensor<16xf32> { ++ // CHECK: "vhlo.dynamic_update_slice_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1<4x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.dynamic_update_slice"(%arg0, %arg1, %arg2) : (tensor<16xf32>, tensor<4xf32>, tensor) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_einsum" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_einsum(%arg0: tensor<8x16xf32>, %arg1: tensor<16x8xf32>) -> tensor<8x8xf32> { ++ // CHECK: "vhlo.einsum_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: einsum_config = #vhlo.string_v1<"ab,bc->ac"> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<8x16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<8x8x!vhlo.f32_v1> ++ %0 = "stablehlo.einsum"(%arg0, %arg1) { ++ einsum_config = "ab,bc->ac" ++ } : (tensor<8x16xf32>, tensor<16x8xf32>) -> tensor<8x8xf32> ++ func.return %0 : tensor<8x8xf32> ++} ++ ++// CHECK-LABEL: "op_exponential_minus_one" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_exponential_minus_one(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.exponential_minus_one_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.exponential_minus_one"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_exponential" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_exponential(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.exponential_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.exponential"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_fft" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_fft(%arg0: tensor<16xcomplex>) -> tensor<16xcomplex> { ++ // CHECK: "vhlo.fft_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: fft_length = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: fft_type = #vhlo ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x!vhlo.complex_v1>) -> !vhlo.tensor_v1<16x!vhlo.complex_v1> ++ %0 = "stablehlo.fft"(%arg0) { ++ fft_type = #stablehlo, ++ fft_length = array ++ } : (tensor<16xcomplex>) -> tensor<16xcomplex> ++ func.return %0 : tensor<16xcomplex> ++} ++ ++// CHECK-LABEL: "op_floor" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_floor(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.floor_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.floor"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++func.func private @op_func(%arg0: tensor {stablehlo.arg = "0"}) -> (tensor {stablehlo.result = "0"}) { ++ // CHECK: "vhlo.func_v1"() <{ ++ // CHECK-SAME: arg_attrs = #vhlo.array_v1<[#vhlo.dict_v1<{#vhlo.string_v1<"stablehlo.arg"> = #vhlo.string_v1<"0">}>]>, ++ // CHECK-SAME: function_type = #vhlo.type_v1) -> !vhlo.tensor_v1>>, ++ // CHECK-SAME: res_attrs = #vhlo.array_v1<[#vhlo.dict_v1<{#vhlo.string_v1<"stablehlo.result"> = #vhlo.string_v1<"0">}>]>, ++ // CHECK-SAME: sym_name = #vhlo.string_v1<"op_func">, ++ // CHECK-SAME: sym_visibility = #vhlo.string_v1<"private"> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG0:.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: "vhlo.return_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : () -> () ++ ++ func.return %arg0 : tensor ++} ++ ++// CHECK-LABEL: "op_gather" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_gather(%arg0 : tensor<2x4x9xf32>, %arg1 : tensor<1x5x2xi32>) -> tensor<1x5x1xf32> { ++ // CHECK: "vhlo.gather_v2"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: collapsed_slice_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: index_vector_dim = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: indices_are_sorted = #vhlo.bool_v1, ++ // CHECK-SAME: offset_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: operand_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>>, ++ // CHECK-SAME: slice_sizes = #vhlo.tensor_v1 : tensor<3xi64>>, ++ // CHECK-SAME: start_index_map = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: start_indices_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<2x4x9x!vhlo.f32_v1>, !vhlo.tensor_v1<1x5x2x!vhlo.i32_v1>) -> !vhlo.tensor_v1<1x5x1x!vhlo.f32_v1> ++ %0 = "stablehlo.gather"(%arg0, %arg1) { ++ dimension_numbers = #stablehlo.gather< ++ offset_dims = [2], ++ collapsed_slice_dims = [0, 1], ++ start_index_map = [0, 1], ++ index_vector_dim = 2 ++ >, ++ slice_sizes = array, ++ indices_are_sorted = true ++ } : (tensor<2x4x9xf32>, tensor<1x5x2xi32>) -> tensor<1x5x1xf32> ++ func.return %0 : tensor<1x5x1xf32> ++} ++ ++// CHECK-LABEL: "op_gather_with_batching_dims" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_gather_with_batching_dims(%arg0 : tensor<5x2x4x9xf32>, %arg1 : tensor<1x5x2xi32>) -> tensor<1x5x1xf32> { ++ // CHECK: "vhlo.gather_v2"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: collapsed_slice_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: index_vector_dim = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: indices_are_sorted = #vhlo.bool_v1, ++ // CHECK-SAME: offset_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: operand_batching_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: slice_sizes = #vhlo.tensor_v1 : tensor<4xi64>>, ++ // CHECK-SAME: start_index_map = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: start_indices_batching_dims = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<5x2x4x9x!vhlo.f32_v1>, !vhlo.tensor_v1<1x5x2x!vhlo.i32_v1>) -> !vhlo.tensor_v1<1x5x1x!vhlo.f32_v1> ++ %0 = "stablehlo.gather"(%arg0, %arg1) { ++ dimension_numbers = #stablehlo.gather< ++ offset_dims = [2], ++ collapsed_slice_dims = [1, 2], ++ operand_batching_dims = [0], ++ start_indices_batching_dims = [1], ++ start_index_map = [1, 2], ++ index_vector_dim = 2 ++ >, ++ slice_sizes = array, ++ indices_are_sorted = true ++ } : (tensor<5x2x4x9xf32>, tensor<1x5x2xi32>) -> tensor<1x5x1xf32> ++ func.return %0 : tensor<1x5x1xf32> ++} ++ ++// CHECK-LABEL: "op_get_dimension_size" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_get_dimension_size(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.get_dimension_size_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: dimension = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.get_dimension_size"(%arg0) { ++ dimension = 0 : i64 ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_get_tuple_element" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_get_tuple_element(%arg0: tuple, tensor>) -> tensor { ++ // CHECK: "vhlo.get_tuple_element_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: index = #vhlo.integer_v1<0 : i32> ++ // CHECK-SAME: }> : (!vhlo.tuple_v1, !vhlo.tensor_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.get_tuple_element"(%arg0) { ++ index = 0 : i32 ++ } : (tuple, tensor>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_if" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_if(%arg0: tensor, %arg1: tensor, %arg2: tensor) -> tensor { ++ // CHECK: "vhlo.if_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: "vhlo.return_v1"(%[[ARG1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }, { ++ // CHECK-NEXT: "vhlo.return_v1"(%[[ARG2]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.if"(%arg0) ({ ++ "stablehlo.return"(%arg1) : (tensor) -> () ++ }, { ++ "stablehlo.return"(%arg2) : (tensor) -> () ++ }) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_imag" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_imag(%arg0: tensor>) -> tensor { ++ // CHECK: "vhlo.imag_v1"(%[[ARG0]]) : (!vhlo.tensor_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.imag"(%arg0) : (tensor>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_infeed" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_infeed(%arg0: !stablehlo.token) -> (tensor, !stablehlo.token) { ++ // CHECK: "vhlo.infeed_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: infeed_config = #vhlo.string_v1<"foo">, ++ // CHECK-SAME{LITERAL}: layout = #vhlo.array_v1<[#vhlo.array_v1<[]>]> ++ // CHECK-SAME: }> : (!vhlo.token_v1) -> (!vhlo.tensor_v1, !vhlo.token_v1) ++ %0:2 = "stablehlo.infeed"(%arg0) { ++ infeed_config = "foo", ++ layout = [[]] ++ } : (!stablehlo.token) -> (tensor, !stablehlo.token) ++ func.return %0#0, %0#1 : tensor, !stablehlo.token ++} ++ ++// CHECK-LABEL: "op_iota" ++func.func @op_iota() -> tensor<16xf32> { ++ // CHECK: "vhlo.iota_v1"() <{ ++ // CHECK-SAME: iota_dimension = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : () -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.iota"() { ++ iota_dimension = 0 : i64 ++ } : () -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_is_finite" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_is_finite(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.is_finite_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.is_finite"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_log" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_log(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.log_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.log"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_log_plus_one" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_log_plus_one(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.log_plus_one_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.log_plus_one"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_logistic" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_logistic(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.logistic_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.logistic"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_map" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_map(%arg0: tensor<16xf32>) -> tensor<16xf32> { ++ // CHECK: "vhlo.map_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: dimensions = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.abs_v1"(%[[ARG1]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.map"(%arg0) ({ ++ ^bb0(%arg1: tensor): ++ %1 = "stablehlo.abs"(%arg1) : (tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ dimensions = array ++ } : (tensor<16xf32>) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_maximum" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_maximum(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.maximum_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.maximum"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_minimum" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_minimum(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.minimum_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.minimum"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_multiply" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_multiply(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.multiply_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.multiply"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_negate" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_negate(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.negate_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.negate"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_not" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_not(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.not_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.not"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_optimization_barrier" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_optimization_barrier(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.optimization_barrier_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.optimization_barrier"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_or" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_or(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.or_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.or"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_outfeed" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_outfeed(%arg0: tensor, %arg1: !stablehlo.token) -> !stablehlo.token { ++ // CHECK: "vhlo.outfeed_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: outfeed_config = #vhlo.string_v1<"foo"> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1, !vhlo.token_v1) -> !vhlo.token_v1 ++ %0 = "stablehlo.outfeed"(%arg0, %arg1) { ++ outfeed_config = "foo" ++ } : (tensor, !stablehlo.token) -> !stablehlo.token ++ func.return %0 : !stablehlo.token ++} ++ ++// CHECK-LABEL: "op_pad" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_pad(%arg0: tensor<8xf32>, %arg1: tensor) -> tensor<16xf32> { ++ // CHECK: "vhlo.pad_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: edge_padding_high = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: edge_padding_low = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: interior_padding = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<8x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.pad"(%arg0, %arg1) { ++ edge_padding_high = array, ++ edge_padding_low = array, ++ interior_padding = array ++ } : (tensor<8xf32>, tensor) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_popcnt" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_popcnt(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.popcnt_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.popcnt"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_power" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_power(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.power_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.power"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_real_dynamic_slice" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}, %[[ARG3:.*]]: {{.*}}) ++func.func @op_real_dynamic_slice(%arg0: tensor, %arg1: tensor<1xindex>, %arg2: tensor<1xindex>, %arg3: tensor<1xindex>) -> tensor { ++ // CHECK: "vhlo.real_dynamic_slice_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]], %[[ARG3]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1<1x!vhlo.index_v1>, !vhlo.tensor_v1<1x!vhlo.index_v1>, !vhlo.tensor_v1<1x!vhlo.index_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.real_dynamic_slice"(%arg0, %arg1, %arg2, %arg3) : (tensor, tensor<1xindex>, tensor<1xindex>, tensor<1xindex>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_real" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_real(%arg0: tensor>) -> tensor { ++ // CHECK: "vhlo.real_v1"(%[[ARG0]]) : (!vhlo.tensor_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.real"(%arg0) : (tensor>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_recv_no_source_target_pairs" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_recv_no_source_target_pairs(%arg0: !stablehlo.token) -> (tensor, !stablehlo.token) { ++ // CHECK: "vhlo.recv_v2"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: channel_type = #vhlo.integer_v1<3 : i64>, ++ // CHECK-SAME: is_host_transfer = #vhlo.bool_v1, ++ // CHECK-SAME{LITERAL}: source_target_pairs = #vhlo.tensor_v1 : tensor<0xi64> ++ // CHECK-SAME{LITERAL}: }> : (!vhlo.token_v1) -> (!vhlo.tensor_v1, !vhlo.token_v1) ++ %0:2 = "stablehlo.recv"(%arg0) { ++ channel_handle = #stablehlo.channel_handle, ++ is_host_transfer = true ++ } : (!stablehlo.token) -> (tensor, !stablehlo.token) ++ func.return %0#0, %0#1 : tensor, !stablehlo.token ++} ++ ++// CHECK-LABEL: "op_reduce" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_reduce(%arg0: tensor<16xf32>, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.reduce_v1"(%[[ARG0]], %[[ARG1]]) ++ // CHECK: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK: "vhlo.return_v1"(%[[VAL1:.*]]) : (!vhlo.tensor_v1) -> () ++ // CHECK: }) : (!vhlo.tensor_v1<16x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.reduce"(%arg0, %arg1) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor): ++ %1 = "stablehlo.add"(%arg2, %arg3) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ dimensions = array ++ } : (tensor<16xf32>, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_reduce_precision" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_reduce_precision(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.reduce_precision_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: exponent_bits = #vhlo.integer_v1<8 : i32> ++ // CHECK-SAME: mantissa_bits = #vhlo.integer_v1<10 : i32> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.reduce_precision"(%arg0) { ++ exponent_bits = 8 : i32, ++ mantissa_bits = 10 : i32 ++ } : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK_lABEL: "op_reduce_with_promotable_types" ++func.func @op_reduce_with_promotable_types(%arg0: tensor<4x4xf32>, %arg1 : tensor) ++ -> (tensor<4xf64>) { ++ // CHECK: "vhlo.reduce_v1"(%[[ARG0:.*]], %[[ARG1:.*]]) ++ // CHECK: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK: "vhlo.return_v1"(%[[VAL1:.*]]) : (!vhlo.tensor_v1) -> () ++ // CHECK: }) : (!vhlo.tensor_v1<4x4x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1<4x!vhlo.f64_v1> ++ %0 = "stablehlo.reduce"(%arg0, %arg1) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor ): ++ %1 = "stablehlo.add"(%arg2, %arg3) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ ++ }) {dimensions = array} : (tensor<4x4xf32>, tensor) -> tensor<4xf64> ++ ++ func.return %0: tensor<4xf64> ++} ++ ++// CHECK-LABEL: "op_reduce_scatter" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_reduce_scatter(%arg0: tensor<16xf32>) -> tensor<16xf32> { ++ // CHECK: "vhlo.reduce_scatter_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<2x1xi64>>, ++ // CHECK-SAME: scatter_dimension = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.add_v1"(%[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.reduce_scatter"(%arg0) ({ ++ ^bb0(%arg1: tensor, %arg2: tensor): ++ %1 = "stablehlo.add"(%arg1, %arg2) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ scatter_dimension = 0 : i64, ++ replica_groups = dense<[[0], [1]]> : tensor<2x1xi64>, ++ channel_handle = #stablehlo.channel_handle, ++ use_global_device_ids ++ } : (tensor<16xf32>) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK_lABEL: "op_reduce_scatter_with_promotable_types" ++func.func @op_reduce_scatter_with_promotable_types(%data: tensor<4x16xf32>) -> tensor<4x4xf64> { ++ // CHECK: "vhlo.reduce_scatter_v1"(%[[ARG0:.*]]) ++ // CHECK: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK: "vhlo.return_v1"(%[[VAL1:.*]]) : (!vhlo.tensor_v1) -> () ++ // CHECK: }) : (!vhlo.tensor_v1<4x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<4x4x!vhlo.f64_v1> ++ %0 = "stablehlo.reduce_scatter"(%data) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor): ++ %1 = stablehlo.add %arg2, %arg3 : tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) {replica_groups = dense<[[0, 1, 2, 3]]> : tensor<1x4xi64>, ++ scatter_dimension = 1 : i64, ++ channel_handle = #stablehlo.channel_handle, ++ use_global_device_ids} : (tensor<4x16xf32>) -> tensor<4x4xf64> ++ func.return %0 : tensor<4x4xf64> ++} ++ ++ ++// CHECK-LABEL: "op_reduce_window" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_reduce_window(%arg0: tensor<2x17x31x7xf32>, %arg1: tensor) -> tensor<2x9x16x7xf32> { ++ // CHECK: "vhlo.reduce_window_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: base_dilations = #vhlo.tensor_v1 : tensor<4xi64>>, ++ // CHECK-SAME{LITERAL}: padding = #vhlo.tensor_v1 : tensor<4x2xi64>>, ++ // CHECK-SAME: window_dilations = #vhlo.tensor_v1 : tensor<4xi64>>, ++ // CHECK-SAME: window_dimensions = #vhlo.tensor_v1 : tensor<4xi64>>, ++ // CHECK-SAME: window_strides = #vhlo.tensor_v1 : tensor<4xi64>> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG2:arg.*]]: !vhlo.tensor_v1, %[[ARG3:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.maximum_v1"(%[[ARG2]], %[[ARG3]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<2x17x31x7x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1<2x9x16x7x!vhlo.f32_v1> ++ %0 = "stablehlo.reduce_window"(%arg0, %arg1) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor): ++ %1 = "stablehlo.maximum"(%arg2, %arg3) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ window_dimensions = array, ++ window_strides = array, ++ base_dilations = array, ++ window_dilations = array, ++ padding = dense<[[0, 0], [2, 0], [0, 2], [0, 0]]> : tensor<4x2xi64> ++ } : (tensor<2x17x31x7xf32>, tensor) -> tensor<2x9x16x7xf32> ++ func.return %0 : tensor<2x9x16x7xf32> ++} ++ ++// CHECK-LABEL: "op_reduce_window_with_promotable_types" ++func.func @op_reduce_window_with_promotable_types(%arg0: tensor<4x2xf32>, ++ %arg1: tensor<4x2xf32>, %init0: tensor, %init1: tensor) -> ++ (tensor<2x2xf64>, tensor<2x2xf32>) { ++ // CHECK: "vhlo.reduce_window_v1"(%[[ARG0:.*]], %[[ARG1:.*]], %[[ARG2:.*]], %[[ARG3:.*]]) ++ // CHECK: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1, %[[ARG3:arg.*]]: !vhlo.tensor_v1, %[[ARG4:arg.*]]: !vhlo.tensor_v1): ++ // CHECK: "vhlo.return_v1"(%[[VAL1:.*]], %[[VAL2:.*]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> () ++ // CHECK: }) : (!vhlo.tensor_v1<4x2x!vhlo.f32_v1>, !vhlo.tensor_v1<4x2x!vhlo.f32_v1>, !vhlo.tensor_v1, !vhlo.tensor_v1) -> (!vhlo.tensor_v1<2x2x!vhlo.f64_v1>, !vhlo.tensor_v1<2x2x!vhlo.f32_v1>) ++ %0:2 = "stablehlo.reduce_window"(%arg0, %arg1, %init0, %init1) ({ ++ ^bb0(%a0: tensor, %a1: tensor, %b0: tensor, ++ %b1: tensor): ++ %2 = stablehlo.add %a0, %b0 : tensor ++ %3 = stablehlo.add %a1, %b1 : tensor ++ "stablehlo.return"(%2,%3) : (tensor, tensor) -> () ++ }) ++ { padding = dense<[[2, 2], [0, 0]]> : tensor<2x2xi64>, ++ window_dimensions = array, ++ window_strides = array } ++ : (tensor<4x2xf32>, tensor<4x2xf32>, tensor, tensor) -> ++ (tensor<2x2xf64>, tensor<2x2xf32>) ++ func.return %0#0, %0#1 : tensor<2x2xf64>, tensor<2x2xf32> ++} ++ ++// CHECK-LABEL: "op_remainder" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_remainder(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.remainder_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.remainder"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_replica_id" ++func.func @op_replica_id() -> tensor { ++ // CHECK: "vhlo.replica_id_v1"() : () -> !vhlo.tensor_v1 ++ %0 = "stablehlo.replica_id"() : () -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_partition_id" ++func.func @op_partition_id() -> tensor { ++ // CHECK: "vhlo.partition_id_v1"() : () -> !vhlo.tensor_v1 ++ %0 = "stablehlo.partition_id"() : () -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_reshape" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_reshape(%arg0: tensor<16xf32>) -> tensor<4x4xf32> { ++ // CHECK: "vhlo.reshape_v1"(%[[ARG0]]) : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<4x4x!vhlo.f32_v1> ++ %0 = "stablehlo.reshape"(%arg0) : (tensor<16xf32>) -> tensor<4x4xf32> ++ func.return %0 : tensor<4x4xf32> ++} ++ ++// CHECK-LABEL: "op_return" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_return(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.case_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: "vhlo.return_v1"(%[[ARG1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.case"(%arg0) ({ ++ "stablehlo.return"(%arg1) : (tensor) -> () ++ }) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_reverse" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_reverse(%arg0: tensor<16xf32>) -> tensor<16xf32> { ++ // CHECK: "vhlo.reverse_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: dimensions = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.reverse"(%arg0) { ++ dimensions = array ++ } : (tensor<16xf32>) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_rng_bit_generator" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_rng_bit_generator(%arg0: tensor) -> (tensor, tensor) { ++ // CHECK: "vhlo.rng_bit_generator_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: rng_algorithm = #vhlo ++ // CHECK-SAME: }> : (!vhlo.tensor_v1) -> (!vhlo.tensor_v1, !vhlo.tensor_v1) ++ %0:2 = "stablehlo.rng_bit_generator"(%arg0) { ++ rng_algorithm = #stablehlo ++ } : (tensor) -> (tensor, tensor) ++ func.return %0#0, %0#1 : tensor, tensor ++} ++ ++// CHECK-LABEL: "op_rng" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_rng(%arg0: tensor, %arg1: tensor, %arg2: tensor<0xindex>) -> tensor { ++ // CHECK: "vhlo.rng_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: rng_distribution = #vhlo ++ // CHECK-SAME: }> : (!vhlo.tensor_v1, !vhlo.tensor_v1, !vhlo.tensor_v1<0x!vhlo.index_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.rng"(%arg0, %arg1, %arg2) { ++ rng_distribution = #stablehlo ++ } : (tensor, tensor, tensor<0xindex>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_round_nearest_afz" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_round_nearest_afz(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.round_nearest_afz_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.round_nearest_afz"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_round_nearest_even" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_round_nearest_even(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.round_nearest_even_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.round_nearest_even"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_rsqrt" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_rsqrt(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.rsqrt_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.rsqrt"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_scatter" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_scatter(%arg0: tensor<200x100x300xf32>, %arg1: tensor<10x2xi32>, %arg2: tensor<10x300xf32>) -> tensor<200x100x300xf32> { ++ // CHECK: "vhlo.scatter_v2"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: index_vector_dim = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: indices_are_sorted = #vhlo.bool_v1, ++ // CHECK-SAME: input_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>>, ++ // CHECK-SAME: inserted_window_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: scatter_dims_to_operand_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: scatter_indices_batching_dims = #vhlo.tensor_v1 : tensor<0xi64>>, ++ // CHECK-SAME: unique_indices = #vhlo.bool_v1, ++ // CHECK-SAME: update_window_dims = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG3:arg.*]]: !vhlo.tensor_v1, %[[ARG4:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.add_v1"(%[[ARG3]], %[[ARG4]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<200x100x300x!vhlo.f32_v1>, !vhlo.tensor_v1<10x2x!vhlo.i32_v1>, !vhlo.tensor_v1<10x300x!vhlo.f32_v1>) -> !vhlo.tensor_v1<200x100x300x!vhlo.f32_v1> ++ %0 = "stablehlo.scatter"(%arg0, %arg1, %arg2) ({ ++ ^bb0(%arg3: tensor, %arg4: tensor): ++ %1 = "stablehlo.add"(%arg3, %arg4) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ scatter_dimension_numbers = #stablehlo.scatter< ++ update_window_dims = [1], ++ inserted_window_dims = [0, 1], ++ scatter_dims_to_operand_dims = [0, 1], ++ index_vector_dim = 1 ++ >, ++ indices_are_sorted = true, ++ unique_indices = true ++ } : (tensor<200x100x300xf32>, tensor<10x2xi32>, tensor<10x300xf32>) -> tensor<200x100x300xf32> ++ func.return %0 : tensor<200x100x300xf32> ++} ++ ++// CHECK-LABEL: "op_scatter_with_batching_dims" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_scatter_with_batching_dims(%arg0: tensor<10x200x100x300xf32>, %arg1: tensor<10x2xi32>, %arg2: tensor<10x300xf32>) -> tensor<10x200x100x300xf32> { ++ // CHECK: "vhlo.scatter_v2"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: index_vector_dim = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: indices_are_sorted = #vhlo.bool_v1, ++ // CHECK-SAME: input_batching_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: inserted_window_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: scatter_dims_to_operand_dims = #vhlo.tensor_v1 : tensor<2xi64>>, ++ // CHECK-SAME: scatter_indices_batching_dims = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: unique_indices = #vhlo.bool_v1, ++ // CHECK-SAME: update_window_dims = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG3:arg.*]]: !vhlo.tensor_v1, %[[ARG4:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.add_v1"(%[[ARG3]], %[[ARG4]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<10x200x100x300x!vhlo.f32_v1>, !vhlo.tensor_v1<10x2x!vhlo.i32_v1>, !vhlo.tensor_v1<10x300x!vhlo.f32_v1>) -> !vhlo.tensor_v1<10x200x100x300x!vhlo.f32_v1> ++ %0 = "stablehlo.scatter"(%arg0, %arg1, %arg2) ({ ++ ^bb0(%arg3: tensor, %arg4: tensor): ++ %1 = "stablehlo.add"(%arg3, %arg4) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ scatter_dimension_numbers = #stablehlo.scatter< ++ update_window_dims = [1], ++ inserted_window_dims = [1, 2], ++ input_batching_dims = [0], ++ scatter_dims_to_operand_dims = [1, 2], ++ scatter_indices_batching_dims = [0], ++ index_vector_dim = 1 ++ >, ++ indices_are_sorted = true, ++ unique_indices = true ++ } : (tensor<10x200x100x300xf32>, tensor<10x2xi32>, tensor<10x300xf32>) -> tensor<10x200x100x300xf32> ++ func.return %0 : tensor<10x200x100x300xf32> ++} ++ ++// CHECK_lABEL: "op_scatter_with_promotable_types" ++func.func @op_scatter_with_promotable_types(%input_tensor: tensor<200x100x300xf32>, ++ %scatter_indices: tensor<10x2xi32>, %updates: tensor<10x300xf32>) -> ++ tensor<200x100x300xf64> { ++ // CHECK: "vhlo.scatter_v2"(%[[ARG0:.*]], %[[ARG1:.*]], %[[ARG2:.*]]) ++ // CHECK: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK: "vhlo.return_v1"(%[[VAL1:.*]]) : (!vhlo.tensor_v1) -> () ++ // CHECK: }) : (!vhlo.tensor_v1<200x100x300x!vhlo.f32_v1>, !vhlo.tensor_v1<10x2x!vhlo.i32_v1>, !vhlo.tensor_v1<10x300x!vhlo.f32_v1>) -> !vhlo.tensor_v1<200x100x300x!vhlo.f64_v1> ++ %0 = "stablehlo.scatter" (%input_tensor, %scatter_indices, %updates) ({ ++ ^bb0(%lhs: tensor, %rhs: tensor): ++ %add = stablehlo.add %lhs, %rhs : tensor ++ "stablehlo.return"(%add) : (tensor) -> () ++ }) { ++ scatter_dimension_numbers = #stablehlo.scatter< ++ update_window_dims = [1], ++ inserted_window_dims = [0, 1], ++ scatter_dims_to_operand_dims = [0, 1], ++ index_vector_dim = 1 ++ >, ++ indices_are_sorted = true, ++ unique_indices = true ++ } : (tensor<200x100x300xf32>, tensor<10x2xi32>, tensor<10x300xf32>) -> ++ tensor<200x100x300xf64> ++ func.return %0 : tensor<200x100x300xf64> ++} ++ ++// CHECK-LABEL: "op_select_and_scatter" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_select_and_scatter(%arg0: tensor<10x24x24x64xf32>, %arg1: tensor<12x13x13x66xf32>, %arg2: tensor) -> tensor<10x24x24x64xf32> { ++ // CHECK: "vhlo.select_and_scatter_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) <{ ++ // CHECK-SAME: padding = #vhlo.tensor_v1 : tensor<4x2xi64>>, ++ // CHECK-SAME: window_dimensions = #vhlo.tensor_v1 : tensor<4xi64>>, ++ // CHECK-SAME: window_strides = #vhlo.tensor_v1 : tensor<4xi64>> ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG31:arg.*]]: !vhlo.tensor_v1, %[[ARG41:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL11:.*]] = "vhlo.compare_v1"(%[[ARG31]], %[[ARG41]]) <{compare_type = #vhlo, comparison_direction = #vhlo}> : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL11]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }, { ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG32:arg.*]]: !vhlo.tensor_v1, %[[ARG42:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL12:.*]] = "vhlo.add_v1"(%[[ARG32]], %[[ARG42]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL12]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<10x24x24x64x!vhlo.f32_v1>, !vhlo.tensor_v1<12x13x13x66x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1<10x24x24x64x!vhlo.f32_v1> ++ %0 = "stablehlo.select_and_scatter"(%arg0, %arg1, %arg2) ({ ++ ^bb0(%arg3: tensor, %arg4: tensor): ++ %1 = "stablehlo.compare"(%arg3, %arg4) {compare_type = #stablehlo, comparison_direction = #stablehlo} : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }, { ++ ^bb0(%arg3: tensor, %arg4: tensor): ++ %1 = "stablehlo.add"(%arg3, %arg4) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ window_dimensions = array, ++ window_strides = array, ++ padding = dense<1> : tensor<4x2xi64> ++ } : (tensor<10x24x24x64xf32>, tensor<12x13x13x66xf32>, tensor) -> tensor<10x24x24x64xf32> ++ func.return %0 : tensor<10x24x24x64xf32> ++} ++ ++// CHECK-LABEL: "op_select_and_scatter_with_promotable_types" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_select_and_scatter_with_promotable_types(%arg0: tensor<10x24x24x64xf32>, %arg1: tensor<12x13x13x66xf32>, %arg2: tensor) -> tensor<10x24x24x64xf64> { ++ // CHECK: "vhlo.select_and_scatter_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) ++ // CHECK: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK: %[[VAL:.*]] = "vhlo.add_v1"(%[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK: "vhlo.return_v1"(%[[VAL]]) : (!vhlo.tensor_v1) -> () ++ // CHECK: }) : (!vhlo.tensor_v1<10x24x24x64x!vhlo.f32_v1>, !vhlo.tensor_v1<12x13x13x66x!vhlo.f32_v1>, !vhlo.tensor_v1) -> !vhlo.tensor_v1<10x24x24x64x!vhlo.f64_v1> ++ %0 = "stablehlo.select_and_scatter"(%arg0, %arg1, %arg2) ({ ++ ^bb0(%arg3: tensor, %arg4: tensor): ++ %1 = "stablehlo.compare"(%arg3, %arg4) {compare_type = #stablehlo, comparison_direction = #stablehlo} : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }, { ++ ^bb0(%arg3: tensor, %arg4: tensor): ++ %1 = "stablehlo.add"(%arg3, %arg4) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ window_dimensions = array, ++ window_strides = array, ++ padding = dense<1> : tensor<4x2xi64> ++ } : (tensor<10x24x24x64xf32>, tensor<12x13x13x66xf32>, tensor) -> tensor<10x24x24x64xf64> ++ func.return %0 : tensor<10x24x24x64xf64> ++} ++ ++// CHECK-LABEL: "op_select" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}, %[[ARG2:.*]]: {{.*}}) ++func.func @op_select(%arg0: tensor, %arg1: tensor, %arg2: tensor) -> tensor { ++ // CHECK: "vhlo.select_v1"(%[[ARG0]], %[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.select"(%arg0, %arg1, %arg2) : (tensor, tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_send_no_source_target_pairs" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_send_no_source_target_pairs(%arg0: tensor, %arg1: !stablehlo.token) -> !stablehlo.token { ++ // CHECK: "vhlo.send_v2"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: channel_type = #vhlo.integer_v1<2 : i64>, ++ // CHECK-SAME: is_host_transfer = #vhlo.bool_v1, ++ // CHECK-SAME{LITERAL}: source_target_pairs = #vhlo.tensor_v1 : tensor<0xi64>> ++ // CHECK-SAME{LITERAL}: }> : (!vhlo.tensor_v1, !vhlo.token_v1) -> !vhlo.token_v1 ++ %0 = "stablehlo.send"(%arg0, %arg1) { ++ channel_handle = #stablehlo.channel_handle, ++ is_host_transfer = true ++ } : (tensor, !stablehlo.token) -> !stablehlo.token ++ func.return %0 : !stablehlo.token ++} ++ ++// CHECK-LABEL: "op_set_dimension_size" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_set_dimension_size(%arg0: tensor, %arg1: tensor) -> tensor<16xf32> { ++ // CHECK: "vhlo.set_dimension_size_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: dimension = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.set_dimension_size"(%arg0, %arg1) { ++ dimension = 0 : i64 ++ } : (tensor, tensor) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_shift_left" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_shift_left(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.shift_left_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.shift_left"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_shift_right_arithmetic" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_shift_right_arithmetic(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.shift_right_arithmetic_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.shift_right_arithmetic"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_shift_right_logical" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_shift_right_logical(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.shift_right_logical_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.shift_right_logical"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_sign" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_sign(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.sign_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.sign"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_sine" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_sine(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.sine_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.sine"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_slice" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_slice(%arg0: tensor<16xf32>) -> tensor<4xf32> { ++ // CHECK: "vhlo.slice_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: limit_indices = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: start_indices = #vhlo.tensor_v1 : tensor<1xi64>>, ++ // CHECK-SAME: strides = #vhlo.tensor_v1 : tensor<1xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<4x!vhlo.f32_v1> ++ %0 = "stablehlo.slice"(%arg0) { ++ start_indices = array, ++ limit_indices = array, ++ strides = array ++ } : (tensor<16xf32>) -> tensor<4xf32> ++ func.return %0 : tensor<4xf32> ++} ++ ++// CHECK-LABEL: "op_sort" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_sort(%arg0: tensor<16xf32>) -> tensor<16xf32> { ++ // CHECK: "vhlo.sort_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: dimension = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: is_stable = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.compare_v1"(%[[ARG1]], %[[ARG2]]) <{compare_type = #vhlo, comparison_direction = #vhlo}> ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x!vhlo.f32_v1> ++ %0 = "stablehlo.sort"(%arg0) ({ ++ ^bb0(%arg1: tensor, %arg2: tensor): ++ %1 = "stablehlo.compare"(%arg1, %arg2) {compare_type = #stablehlo, comparison_direction = #stablehlo} : (tensor, tensor) -> tensor ++ "stablehlo.return"(%1) : (tensor) -> () ++ }) { ++ dimension = 0 : i64, ++ is_stable = true ++ } : (tensor<16xf32>) -> tensor<16xf32> ++ func.return %0 : tensor<16xf32> ++} ++ ++// CHECK-LABEL: "op_sqrt" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_sqrt(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.sqrt_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.sqrt"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_subtract" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_subtract(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.subtract_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.subtract"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_tan" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_tan(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.tan_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.tan"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_tanh" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_tanh(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.tanh_v2"(%[[ARG0]]) <{result_accuracy = #vhlo.result_accuracy_v1>}> : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.tanh"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_torch_index_select" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_torch_index_select(%arg0: tensor<5x1x5xf32>, %arg1: tensor<2xi32>) -> tensor<2x1x5xf32> { ++ // CHECK: "vhlo.torch_index_select_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: batch_dims = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: dim = #vhlo.integer_v1<0 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<5x1x5x!vhlo.f32_v1>, !vhlo.tensor_v1<2x!vhlo.i32_v1>) -> !vhlo.tensor_v1<2x1x5x!vhlo.f32_v1> ++ %0 = "stablehlo.torch_index_select"(%arg0, %arg1) { ++ dim = 0 : i64, ++ batch_dims = 0 : i64 ++ } : (tensor<5x1x5xf32>, tensor<2xi32>) -> tensor<2x1x5xf32> ++ func.return %0 : tensor<2x1x5xf32> ++} ++ ++// CHECK-LABEL: "op_transpose" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_transpose(%arg0: tensor<16x8xf32>) -> tensor<8x16xf32> { ++ // CHECK: "vhlo.transpose_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: permutation = #vhlo.tensor_v1 : tensor<2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<8x16x!vhlo.f32_v1> ++ %0 = "stablehlo.transpose"(%arg0) { ++ permutation = array ++ } : (tensor<16x8xf32>) -> tensor<8x16xf32> ++ func.return %0 : tensor<8x16xf32> ++} ++ ++// CHECK-LABEL: "op_triangular_solve" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_triangular_solve(%arg0: tensor<16x16xf32>, %arg1: tensor<16x16xf32>) -> tensor<16x16xf32> { ++ // CHECK: "vhlo.triangular_solve_v1"(%[[ARG0]], %[[ARG1]]) <{ ++ // CHECK-SAME: left_side = #vhlo.bool_v1, ++ // CHECK-SAME: lower = #vhlo.bool_v1, ++ // CHECK-SAME: transpose_a = #vhlo, ++ // CHECK-SAME: unit_diagonal = #vhlo.bool_v1 ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x16x!vhlo.f32_v1>, !vhlo.tensor_v1<16x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x16x!vhlo.f32_v1> ++ %0 = "stablehlo.triangular_solve"(%arg0, %arg1) { ++ left_side = true, ++ lower = true, ++ unit_diagonal = true, ++ transpose_a = #stablehlo ++ } : (tensor<16x16xf32>, tensor<16x16xf32>) -> tensor<16x16xf32> ++ func.return %0 : tensor<16x16xf32> ++} ++ ++// CHECK-LABEL: "op_tuple" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_tuple(%arg0: tensor) -> tuple> { ++ // CHECK: "vhlo.tuple_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tuple_v1> ++ %0 = "stablehlo.tuple"(%arg0) : (tensor) -> tuple> ++ func.return %0 : tuple> ++} ++ ++// CHECK-LABEL: "op_unary_einsum" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_unary_einsum(%arg0: tensor<8x16xf32>) -> tensor<8xf32> { ++ // CHECK: "vhlo.unary_einsum_v1"(%[[ARG0]]) <{ ++ // CHECK-SAME: einsum_config = #vhlo.string_v1<"ab->a"> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<8x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<8x!vhlo.f32_v1> ++ %0 = "stablehlo.unary_einsum"(%arg0) { ++ einsum_config = "ab->a" ++ } : (tensor<8x16xf32>) -> tensor<8xf32> ++ func.return %0 : tensor<8xf32> ++} ++ ++// CHECK-LABEL: "op_uniform_dequantize" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_uniform_dequantize(%arg0: tensor>) -> tensor { ++ // CHECK: "vhlo.uniform_dequantize_v1"(%[[ARG0]]) : (!vhlo.tensor_v1>) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.uniform_dequantize"(%arg0) : (tensor>) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "op_uniform_quantize" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_uniform_quantize(%arg0: tensor) -> tensor> { ++ // CHECK: "vhlo.uniform_quantize_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1> ++ %0 = "stablehlo.uniform_quantize"(%arg0) : (tensor) -> tensor> ++ func.return %0 : tensor> ++} ++ ++// CHECK-LABEL: "op_while" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_while(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.while_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: "vhlo.return_v1"(%[[ARG1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }, { ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1) ++ // CHECK-NEXT: "vhlo.return_v1"(%[[ARG1]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.while"(%arg0) ({ ++ ^bb0(%arg1: tensor): ++ "stablehlo.return"(%arg1) : (tensor) -> () ++ }, { ++ ^bb0(%arg1: tensor): ++ "stablehlo.return"(%arg1) : (tensor) -> () ++ }) : (tensor) -> tensor ++ func.return %0: tensor ++} ++ ++// CHECK-LABEL: "op_xor" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @op_xor(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.xor_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.xor"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// ============ TYPES ============ ++ ++// CHECK-LABEL: "type_i1" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_i1(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.and_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.and"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_i2" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_i2(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_i4" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_i4(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_i8" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_i8(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_i16" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_i16(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_i32" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_i32(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_i64" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_i64(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_ui2" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_ui2(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_ui4" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_ui4(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_ui8" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_ui8(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_ui16" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_ui16(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_ui32" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_ui32(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_ui64" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_ui64(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f4E2M1FN" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f4E2M1FN(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f6E2M3FN" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f6E2M3FN(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f6E3M2FN" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f6E3M2FN(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f8E3M4" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f8E3M4(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f8E4M3" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f8E4M3(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f8E4M3FN" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f8E4M3FN(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f8E5M2" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f8E5M2(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f8E4M3FNUZ" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f8E4M3FNUZ(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f8E4M3B11FNUZ" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f8E4M3B11FNUZ(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f8E5M2FNUZ" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f8E5M2FNUZ(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f8E8M0FNU" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f8E8M0FNU(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_bf16" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_bf16(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f16" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f16(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f32" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f32(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_f64" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_f64(%arg0: tensor, %arg1: tensor) -> tensor { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor, tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_complex_f32" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_complex_f32(%arg0: tensor>, %arg1: tensor>) -> tensor> { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1>, !vhlo.tensor_v1>) -> !vhlo.tensor_v1> ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor>, tensor>) -> tensor> ++ func.return %0 : tensor> ++} ++ ++// CHECK-LABEL: "type_complex_f64" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_complex_f64(%arg0: tensor>, %arg1: tensor>) -> tensor> { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1>, !vhlo.tensor_v1>) -> !vhlo.tensor_v1> ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor>, tensor>) -> tensor> ++ func.return %0 : tensor> ++} ++ ++// CHECK-LABEL: "type_tf32" ++// CHECK: #vhlo.type_v1 ++func.func @type_tf32() attributes {stablehlo.attr = tf32 } { ++ return ++} ++ ++// CHECK-LABEL: "type_none" ++// CHECK: #vhlo.type_v1 ++func.func @type_none() attributes {stablehlo.attr = none } { ++ return ++} ++ ++// CHECK-LABEL: "type_dynamism_ranked" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @type_dynamism_ranked(%arg0: tensor) -> tensor { ++ // CHECK: "vhlo.abs_v1"(%[[ARG0]]) : (!vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ %0 = "stablehlo.abs"(%arg0) : (tensor) -> tensor ++ func.return %0 : tensor ++} ++ ++// CHECK-LABEL: "type_per_tensor_quantization" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @type_per_tensor_quantization(%arg0: tensor>, %arg1: tensor>) -> tensor> { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG1]]) : (!vhlo.tensor_v1>, !vhlo.tensor_v1>) -> !vhlo.tensor_v1> ++ %0 = "stablehlo.add"(%arg0, %arg1) : (tensor>, tensor>) -> tensor> ++ func.return %0 : tensor> ++} ++ ++// CHECK-LABEL: "type_per_axis_quantization" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @type_per_axis_quantization(%arg0: tensor<2x!quant.uniform>) -> tensor<2x!quant.uniform> { ++ // CHECK: "vhlo.add_v1"(%[[ARG0]], %[[ARG0]]) : (!vhlo.tensor_v1<2x!vhlo.quant_per_axis_v1>, !vhlo.tensor_v1<2x!vhlo.quant_per_axis_v1>) -> !vhlo.tensor_v1<2x!vhlo.quant_per_axis_v1> ++ %0 = stablehlo.add %arg0, %arg0 : tensor<2x!quant.uniform> ++ func.return %0 : tensor<2x!quant.uniform> ++} ++ ++// CHECK: function_type = #vhlo.type_v1 !vhlo.token_v1>> ++// CHECK-LABEL: "type_token_callee" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @type_token_callee(%arg0: !stablehlo.token) -> !stablehlo.token { ++ // CHECK: "vhlo.return_v1"(%[[ARG0]]) : (!vhlo.token_v1) -> () ++ return %arg0 : !stablehlo.token ++} ++ ++// CHECK: function_type = #vhlo.type_v1 !vhlo.token_v1>> ++// CHECK-LABEL: "type_token_caller" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @type_token_caller(%arg0: !stablehlo.token) -> !stablehlo.token { ++ // CHECK: "vhlo.call_v1"(%[[ARG0]]) <{callee = #vhlo.string_v1<"type_token_callee">} ++ // CHECK-SAME: (!vhlo.token_v1) -> !vhlo.token_v1 ++ %0 = func.call @type_token_callee(%arg0) : (!stablehlo.token) -> !stablehlo.token ++ return %0 : !stablehlo.token ++} ++ ++// CHECK-LABEL: "type_tuple" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @type_tuple(%arg0: tuple>) -> tuple { ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "foo" ++ // CHECK: (!vhlo.tuple_v1>) -> !vhlo.tuple_v1 ++ } : (tuple>) -> tuple ++ return %0 : tuple ++} ++ ++// CHECK-LABEL: type_buffer_function_input_output ++// CHECK-NEXT: (%[[ARG0:.*]]: !vhlo.buffer_v1<2x!vhlo.f32_v1>) ++func.func @type_buffer_function_input_output(%arg0: memref<2xf32>) -> memref<2xf32> { ++ // CHECK: "vhlo.return_v1"(%[[ARG0]]) : (!vhlo.buffer_v1<2x!vhlo.f32_v1>) -> () ++ func.return %arg0 : memref<2xf32> ++} ++ ++// CHECK-LABEL: type_buffer_special_custom_calls ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @type_buffer_special_custom_calls(%arg0: tensor<2xf32>) -> tensor<2xf32> { ++ // CHECK: %[[CALL0:.*]] = "vhlo.custom_call_v2"(%[[ARG0]]) ++ // CHECK-SAME: call_target_name = #vhlo.string_v1<"Pin"> ++ // CHECK-SAME: : (!vhlo.tensor_v1<2x!vhlo.f32_v1>) -> !vhlo.buffer_v1<2x!vhlo.f32_v1> ++ %0 = "stablehlo.custom_call"(%arg0) { ++ call_target_name = "Pin", ++ api_version = 4 : i32 ++ } : (tensor<2xf32>) -> memref<2xf32> ++ // CHECK: %{{.*}} = "vhlo.custom_call_v2"(%[[CALL0]]) ++ // CHECK-SAME: call_target_name = #vhlo.string_v1<"Unpin"> ++ // CHECK-SAME: : (!vhlo.buffer_v1<2x!vhlo.f32_v1>) -> !vhlo.tensor_v1<2x!vhlo.f32_v1> ++ %1 = "stablehlo.custom_call"(%0) { ++ call_target_name = "Unpin", ++ api_version = 4 : i32 ++ } : (memref<2xf32>) -> tensor<2xf32> ++ func.return %1 : tensor<2xf32> ++} ++ ++// CHECK: function_type = #vhlo.type_v1 !vhlo.future_v1>>> ++// CHECK-LABEL: "type_future" ++func.func private @type_future() -> !stablehlo.future> ++ ++// CHECK: function_type = #vhlo.type_v1 !vhlo.future_v1>>> ++// CHECK-LABEL: "type_future_ranked" ++func.func private @type_future_ranked() -> !stablehlo.future> ++ ++// CHECK: function_type = #vhlo.type_v1 !vhlo.future_v1>>> ++// CHECK-LABEL: "type_future_dynamic" ++func.func private @type_future_dynamic() -> !stablehlo.future> ++ ++// CHECK: function_type = #vhlo.type_v1 !vhlo.future_v1, !vhlo.tensor_v1>>> ++// CHECK-LABEL: "type_future_multiple" ++func.func private @type_future_multiple() -> !stablehlo.future, tensor> ++ ++// CHECK-LABEL: "op_async_start_all_reduce" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_async_start_all_reduce(%arg0: tensor<4x4xf32>) -> !stablehlo.future> { ++ // CHECK: "vhlo.async_start_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[BARG0:.*]]: {{.*}}): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.all_reduce_v2"(%[[BARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<1x2xi64>>, ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL2:.*]] = "vhlo.add_v1"(%[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL2]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<4x4x!vhlo.f32_v1>) -> !vhlo.tensor_v1<4x4x!vhlo.f32_v1> ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1<4x4x!vhlo.f32_v1>) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<4x4x!vhlo.f32_v1>) -> !vhlo.future_v1> ++ %0 = "stablehlo.async_start"(%arg0) ({ ++ ^bb0(%barg0: tensor<4x4xf32>): ++ %1 = "stablehlo.all_reduce"(%barg0) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor): ++ %2 = "stablehlo.add"(%arg2, %arg3) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%2) : (tensor) -> () ++ }) {replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>} : (tensor<4x4xf32>) -> tensor<4x4xf32> ++ "stablehlo.return"(%1) : (tensor<4x4xf32>) -> () ++ }) : (tensor<4x4xf32>) -> !stablehlo.future> ++ func.return %0: !stablehlo.future> ++} ++ ++// CHECK-LABEL: "op_async_start_all_gather" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_async_start_all_gather(%arg0: tensor<8x2xf32>) -> !stablehlo.future> { ++ // CHECK: "vhlo.async_start_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[BARG0:.*]]: {{.*}}): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.all_gather_v2"(%[[BARG0]]) <{ ++ // CHECK-SAME: all_gather_dim = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<2x4xi64>>, ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<8x2x!vhlo.f32_v1>) -> !vhlo.tensor_v1<8x8x!vhlo.f32_v1> ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1<8x8x!vhlo.f32_v1>) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<8x2x!vhlo.f32_v1>) -> !vhlo.future_v1> ++ %0 = "stablehlo.async_start"(%arg0) ({ ++ ^bb0(%barg0: tensor<8x2xf32>): ++ %1 = "stablehlo.all_gather"(%barg0) { ++ all_gather_dim = 1 : i64, ++ replica_groups = dense<[[0, 2, 4, 6], [1, 3, 5, 7]]> : tensor<2x4xi64> ++ } : (tensor<8x2xf32>) -> tensor<8x8xf32> ++ "stablehlo.return"(%1) : (tensor<8x8xf32>) -> () ++ }) : (tensor<8x2xf32>) -> !stablehlo.future> ++ func.return %0: !stablehlo.future> ++} ++ ++// CHECK-LABEL: "op_async_start_all_to_all" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_async_start_all_to_all(%arg0: tensor<4x16xf32>) -> !stablehlo.future> { ++ // CHECK: "vhlo.async_start_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[BARG0:.*]]: {{.*}}): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.all_to_all_v2"(%[[BARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: concat_dimension = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<1x4xi64>>, ++ // CHECK-SAME: split_count = #vhlo.integer_v1<4 : i64>, ++ // CHECK-SAME: split_dimension = #vhlo.integer_v1<1 : i64> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<4x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x4x!vhlo.f32_v1> ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1<16x4x!vhlo.f32_v1>) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<4x16x!vhlo.f32_v1>) -> !vhlo.future_v1> ++ %0 = "stablehlo.async_start"(%arg0) ({ ++ ^bb0(%barg0: tensor<4x16xf32>): ++ %1 = "stablehlo.all_to_all"(%barg0) { ++ split_dimension = 1 : i64, ++ concat_dimension = 0 : i64, ++ split_count = 4 : i64, ++ replica_groups = dense<[[0, 1, 2, 3]]> : tensor<1x4xi64> ++ } : (tensor<4x16xf32>) -> tensor<16x4xf32> ++ "stablehlo.return"(%1) : (tensor<16x4xf32>) -> () ++ }) : (tensor<4x16xf32>) -> !stablehlo.future> ++ func.return %0: !stablehlo.future> ++} ++ ++// CHECK-LABEL: "op_async_start_collective_permute" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_async_start_collective_permute(%arg0: tensor<128x32xf32>) -> !stablehlo.future> { ++ // CHECK: "vhlo.async_start_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[BARG0:.*]]: {{.*}}): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.collective_permute_v1"(%[[BARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: source_target_pairs = #vhlo.tensor_v1 : tensor<3x2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<128x32x!vhlo.f32_v1>) -> !vhlo.tensor_v1<128x32x!vhlo.f32_v1> ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1<128x32x!vhlo.f32_v1>) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<128x32x!vhlo.f32_v1>) -> !vhlo.future_v1> ++ %0 = "stablehlo.async_start"(%arg0) ({ ++ ^bb0(%barg0: tensor<128x32xf32>): ++ %1 = "stablehlo.collective_permute"(%barg0) { ++ source_target_pairs = dense<[[0, 1], [1, 2], [2, 3]]> : tensor<3x2xi64> ++ } : (tensor<128x32xf32>) -> tensor<128x32xf32> ++ "stablehlo.return"(%1) : (tensor<128x32xf32>) -> () ++ }) : (tensor<128x32xf32>) -> !stablehlo.future> ++ func.return %0: !stablehlo.future> ++} ++ ++// CHECK-LABEL: "op_async_start_collective_broadcast" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_async_start_collective_broadcast(%arg0: tensor<16x8xf32>) -> !stablehlo.future> { ++ // CHECK: "vhlo.async_start_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[BARG0:.*]]: {{.*}}): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.collective_broadcast_v2"(%[[BARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME: has_dynamic_root = #vhlo.bool_v1, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<1x2xi64>> ++ // CHECK-SAME: }> : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.tensor_v1<16x8x!vhlo.f32_v1> ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<16x8x!vhlo.f32_v1>) -> !vhlo.future_v1> ++ %0 = "stablehlo.async_start"(%arg0) ({ ++ ^bb0(%barg0: tensor<16x8xf32>): ++ %1 = "stablehlo.collective_broadcast"(%barg0) { ++ replica_groups = dense<[[0, 1]]> : tensor<1x2xi64> ++ } : (tensor<16x8xf32>) -> tensor<16x8xf32> ++ "stablehlo.return"(%1) : (tensor<16x8xf32>) -> () ++ }) : (tensor<16x8xf32>) -> !stablehlo.future> ++ func.return %0: !stablehlo.future> ++} ++ ++// CHECK-LABEL: "op_async_start_reduce_scatter" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_async_start_reduce_scatter(%arg0: tensor<4x16xf32>) -> !stablehlo.future> { ++ // CHECK: "vhlo.async_start_v1"(%[[ARG0]]) ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[BARG0:.*]]: {{.*}}): ++ // CHECK-NEXT: %[[VAL1:.*]] = "vhlo.reduce_scatter_v1"(%[[BARG0]]) <{ ++ // CHECK-SAME: channel_id = #vhlo.integer_v1<0 : i64>, ++ // CHECK-SAME{LITERAL}: replica_groups = #vhlo.tensor_v1 : tensor<1x4xi64>>, ++ // CHECK-SAME: scatter_dimension = #vhlo.integer_v1<1 : i64>, ++ // CHECK-SAME: use_global_device_ids = #vhlo.bool_v1 ++ // CHECK-SAME: }> ({ ++ // CHECK-NEXT: ^[[BB:bb.*]](%[[ARG1:arg.*]]: !vhlo.tensor_v1, %[[ARG2:arg.*]]: !vhlo.tensor_v1): ++ // CHECK-NEXT: %[[VAL2:.*]] = "vhlo.add_v1"(%[[ARG1]], %[[ARG2]]) : (!vhlo.tensor_v1, !vhlo.tensor_v1) -> !vhlo.tensor_v1 ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL2]]) : (!vhlo.tensor_v1) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<4x16x!vhlo.f32_v1>) -> !vhlo.tensor_v1<4x4x!vhlo.f32_v1> ++ // CHECK-NEXT: "vhlo.return_v1"(%[[VAL1]]) : (!vhlo.tensor_v1<4x4x!vhlo.f32_v1>) -> () ++ // CHECK-NEXT: }) : (!vhlo.tensor_v1<4x16x!vhlo.f32_v1>) -> !vhlo.future_v1> ++ %0 = "stablehlo.async_start"(%arg0) ({ ++ ^bb0(%barg0: tensor<4x16xf32>): ++ %1 = "stablehlo.reduce_scatter"(%barg0) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor): ++ %2 = stablehlo.add %arg2, %arg3 : tensor ++ "stablehlo.return"(%2) : (tensor) -> () ++ }) {replica_groups = dense<[[0, 1, 2, 3]]> : tensor<1x4xi64>, ++ scatter_dimension = 1 : i64} : (tensor<4x16xf32>) -> tensor<4x4xf32> ++ "stablehlo.return"(%1) : (tensor<4x4xf32>) -> () ++ }) : (tensor<4x16xf32>) -> !stablehlo.future> ++ func.return %0: !stablehlo.future> ++} ++ ++// CHECK-LABEL: "op_async_done" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}) ++func.func @op_async_done(%arg0: tensor<4x4xf32>) -> tensor<4x4xf32> { ++ // CHECK: %[[VAL1:.*]] = "vhlo.async_start_v1"(%[[ARG0]]) ++ // CHECK: %[[VAL2:.*]] = "vhlo.async_done_v1"(%[[VAL1]]) : (!vhlo.future_v1>) -> !vhlo.tensor_v1<4x4x!vhlo.f32_v1> ++ %0 = "stablehlo.async_start"(%arg0) ({ ++ ^bb0(%barg0: tensor<4x4xf32>): ++ %1 = "stablehlo.all_reduce"(%barg0) ({ ++ ^bb0(%arg2: tensor, %arg3: tensor): ++ %2 = "stablehlo.add"(%arg2, %arg3) : (tensor, tensor) -> tensor ++ "stablehlo.return"(%2) : (tensor) -> () ++ }) {replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>} : (tensor<4x4xf32>) -> tensor<4x4xf32> ++ "stablehlo.return"(%1) : (tensor<4x4xf32>) -> () ++ }) : (tensor<4x4xf32>) -> !stablehlo.future> ++ %1 = "stablehlo.async_done"(%0) : (!stablehlo.future>) -> tensor<4x4xf32> ++ func.return %1: tensor<4x4xf32> ++} ++ ++// ============ DEPENDENCIES ============ ++ ++func.func @composite_target(%arg0: tensor) -> tensor { ++ return %arg0: tensor ++} +diff --ruN a/stablehlo/stablehlo/tests/vhlo/stablehlo_legalize_to_vhlo.mlir b/stablehlo/stablehlo/tests/vhlo/stablehlo_legalize_to_vhlo.mlir +--- stablehlo/stablehlo/tests/vhlo/stablehlo_legalize_to_vhlo.mlir ++++ stablehlo/stablehlo/tests/vhlo/stablehlo_legalize_to_vhlo.mlir +@@ -749,6 +749,80 @@ + func.return %0 : tensor<8x8x8xf32> + } + ++// CHECK-LABEL: "dot_general_algorithm_f8e4m3fn_x3" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @dot_general_algorithm_f8e4m3fn_x3(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { ++// CHECK: "vhlo.dot_general_v2"(%[[ARG0]], %[[ARG1]]) <{ ++// CHECK-SAME: accumulation_type = #vhlo.type_v1, ++// CHECK-SAME: allow_imprecise_accumulation = #vhlo.bool_v1, ++// CHECK-SAME: lhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: lhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_precision_type = #vhlo.type_v1, ++// CHECK-SAME: num_primitive_operations = #vhlo.integer_v1<3 : i64>, ++// CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++// CHECK-SAME: rhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: rhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_precision_type = #vhlo.type_v1 ++// CHECK-SAME: }> : (!vhlo.tensor_v1<8x8x16x!vhlo.bf16_v1>, !vhlo.tensor_v1<8x16x8x!vhlo.bf16_v1>) -> !vhlo.tensor_v1<8x8x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) { ++ dot_dimension_numbers = #stablehlo.dot< ++ lhs_batching_dimensions = [0], ++ lhs_contracting_dimensions = [2], ++ rhs_batching_dimensions = [0], ++ rhs_contracting_dimensions = [1] ++ >, ++ algorithm = #stablehlo.dot_algorithm< ++ lhs_precision_type = f8E4M3FN, ++ rhs_precision_type = f8E4M3FN, ++ accumulation_type = f32, ++ lhs_component_count = 1, ++ rhs_component_count = 1, ++ num_primitive_operations = 3, ++ allow_imprecise_accumulation = false ++ > ++ } : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> ++ func.return %0 : tensor<8x8x8xf32> ++} ++ ++// CHECK-LABEL: "dot_general_algorithm_f8e4m3fn_x4" ++// CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) ++func.func @dot_general_algorithm_f8e4m3fn_x4(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { ++// CHECK: "vhlo.dot_general_v2"(%[[ARG0]], %[[ARG1]]) <{ ++// CHECK-SAME: accumulation_type = #vhlo.type_v1, ++// CHECK-SAME: allow_imprecise_accumulation = #vhlo.bool_v1, ++// CHECK-SAME: lhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: lhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: lhs_precision_type = #vhlo.type_v1, ++// CHECK-SAME: num_primitive_operations = #vhlo.integer_v1<4 : i64>, ++// CHECK-SAME: precision_config = #vhlo.array_v1<[#vhlo, #vhlo]>, ++// CHECK-SAME: rhs_batching_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_component_count = #vhlo.integer_v1<1 : i64>, ++// CHECK-SAME: rhs_contracting_dimensions = #vhlo.tensor_v1 : tensor<1xi64>>, ++// CHECK-SAME: rhs_precision_type = #vhlo.type_v1 ++// CHECK-SAME: }> : (!vhlo.tensor_v1<8x8x16x!vhlo.bf16_v1>, !vhlo.tensor_v1<8x16x8x!vhlo.bf16_v1>) -> !vhlo.tensor_v1<8x8x8x!vhlo.f32_v1> ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) { ++ dot_dimension_numbers = #stablehlo.dot< ++ lhs_batching_dimensions = [0], ++ lhs_contracting_dimensions = [2], ++ rhs_batching_dimensions = [0], ++ rhs_contracting_dimensions = [1] ++ >, ++ algorithm = #stablehlo.dot_algorithm< ++ lhs_precision_type = f8E4M3FN, ++ rhs_precision_type = f8E4M3FN, ++ accumulation_type = f32, ++ lhs_component_count = 1, ++ rhs_component_count = 1, ++ num_primitive_operations = 4, ++ allow_imprecise_accumulation = false ++ > ++ } : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> ++ func.return %0 : tensor<8x8x8xf32> ++} ++ + // CHECK-LABEL: "default_dynamic_broadcast_in_dim" + // CHECK-NEXT: (%[[ARG0:.*]]: {{.*}}, %[[ARG1:.*]]: {{.*}}) + func.func @default_dynamic_broadcast_in_dim(%arg0: tensor, %arg1: tensor<2xindex>) -> tensor { +diff --ruN a/stablehlo/stablehlo/tests/vhlo/vhlo_to_version_downgrade_invalid.1_20_0.mlir b/stablehlo/stablehlo/tests/vhlo/vhlo_to_version_downgrade_invalid.1_20_0.mlir +--- stablehlo/stablehlo/tests/vhlo/vhlo_to_version_downgrade_invalid.1_20_0.mlir ++++ stablehlo/stablehlo/tests/vhlo/vhlo_to_version_downgrade_invalid.1_20_0.mlir +@@ -0,0 +1,45 @@ ++// RUN: stablehlo-opt --stablehlo-legalize-to-vhlo --vhlo-to-version='target=1.20.0' --verify-diagnostics --split-input-file %s ++ ++// expected-error @+1 {{failed to convert VHLO to v1.20.0}} ++module { ++ func.func @dot_general_algorithm_f8e4m3fn_x3(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { ++ // expected-error @+1 {{failed to legalize operation 'vhlo.dot_general_v2' that was explicitly marked illegal}} ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) <{ ++ dot_dimension_numbers = #stablehlo.dot, ++ precision_config = [#stablehlo, #stablehlo], ++ algorithm = #stablehlo.dot_algorithm< ++ lhs_precision_type = f8E4M3FN, ++ rhs_precision_type = f8E4M3FN, ++ accumulation_type = f32, ++ lhs_component_count = 1, ++ rhs_component_count = 1, ++ num_primitive_operations = 3, ++ allow_imprecise_accumulation = false ++ > ++ }> : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> ++ func.return %0 : tensor<8x8x8xf32> ++ } ++} ++ ++// ----- ++ ++// expected-error @+1 {{failed to convert VHLO to v1.20.0}} ++module { ++ func.func @dot_general_algorithm_f8e4m3fn_x4(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { ++ // expected-error @+1 {{failed to legalize operation 'vhlo.dot_general_v2' that was explicitly marked illegal}} ++ %0 = "stablehlo.dot_general"(%arg0, %arg1) <{ ++ dot_dimension_numbers = #stablehlo.dot, ++ precision_config = [#stablehlo, #stablehlo], ++ algorithm = #stablehlo.dot_algorithm< ++ lhs_precision_type = f8E4M3FN, ++ rhs_precision_type = f8E4M3FN, ++ accumulation_type = f32, ++ lhs_component_count = 1, ++ rhs_component_count = 1, ++ num_primitive_operations = 4, ++ allow_imprecise_accumulation = false ++ > ++ }> : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> ++ func.return %0 : tensor<8x8x8xf32> ++ } ++} diff --ruN a/stablehlo/stablehlo/tools/StablehloLspServerMain.cpp b/stablehlo/stablehlo/tools/StablehloLspServerMain.cpp --- stablehlo/stablehlo/tools/StablehloLspServerMain.cpp +++ stablehlo/stablehlo/tools/StablehloLspServerMain.cpp diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/tests/fusion_emitter_device_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/tests/fusion_emitter_device_test.cc index d4c17e5f31c651..22f2360d567de0 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/tests/fusion_emitter_device_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/tests/fusion_emitter_device_test.cc @@ -1922,6 +1922,9 @@ ErrorSpec ErrorSpecForDotAlgorithm(PrecisionConfig::Algorithm algorithm) { case PrecisionConfig::ALG_DOT_ANY_F8_ANY_F8_F32: case PrecisionConfig::ALG_DOT_ANY_F8_ANY_F8_F32_FAST_ACCUM: return kExactMatch; + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X3: + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X4: + return default_error_spec; // Keep in order to make the switch exhaustive. case PrecisionConfig_Algorithm_PrecisionConfig_Algorithm_INT_MIN_SENTINEL_DO_NOT_USE_: // NOLINT(whitespace/line_length) case PrecisionConfig_Algorithm_PrecisionConfig_Algorithm_INT_MAX_SENTINEL_DO_NOT_USE_: // NOLINT(whitespace/line_length) diff --git a/third_party/xla/xla/hlo/translate/hlo_to_mhlo/attribute_importer.cc b/third_party/xla/xla/hlo/translate/hlo_to_mhlo/attribute_importer.cc index a7fe71c9520948..e2632e1a13ac1d 100644 --- a/third_party/xla/xla/hlo/translate/hlo_to_mhlo/attribute_importer.cc +++ b/third_party/xla/xla/hlo/translate/hlo_to_mhlo/attribute_importer.cc @@ -231,6 +231,18 @@ mlir::stablehlo::DotAlgorithmAttr ConvertDotAlgorithm( lhs = rhs = accum = builder->getF64Type(); break; } + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X3: { + lhs = rhs = builder->getType(); + accum = builder->getF32Type(); + numPrimitiveOperations = 3; + break; + } + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X4: { + lhs = rhs = builder->getType(); + accum = builder->getF32Type(); + numPrimitiveOperations = 4; + break; + } default: // Unset, sentinels return mlir::stablehlo::DotAlgorithmAttr{}; diff --git a/third_party/xla/xla/hlo/translate/hlo_to_mhlo/tests/attributes.hlo b/third_party/xla/xla/hlo/translate/hlo_to_mhlo/tests/attributes.hlo index 7cc92d0ce8d2f8..99f017a5bcdffb 100644 --- a/third_party/xla/xla/hlo/translate/hlo_to_mhlo/tests/attributes.hlo +++ b/third_party/xla/xla/hlo/translate/hlo_to_mhlo/tests/attributes.hlo @@ -142,3 +142,27 @@ ENTRY %main.4 (Arg_0.1: f64[2,2,2], Arg_1.2: f64[2,2,2]) -> f64[2,2,2] { %Arg_1.2 = f64[2,2,2] parameter(1) ROOT %dot.3 = f64[2,2,2] dot(f64[2,2,2] %Arg_0.1, f64[2,2,2] %Arg_1.2), lhs_batch_dims={0}, lhs_contracting_dims={2}, rhs_batch_dims={0}, rhs_contracting_dims={1}, algorithm=dot_f64_f64_f64, metadata={source_file="within split at third_party/tensorflow/compiler/xla/hlo/translate/mhlo_to_hlo/tests/attributes.mlir:243 offset " source_line=7} } + +// ----- + +HloModule dot_algorithm_bf16_bf16_fp8x3, entry_computation_layout={(bf16[2,2,2]{2,1,0}, bf16[2,2,2]{2,1,0})->f32[2,2,2]{2,1,0}} + +// CHECK-LABEL: module @dot_algorithm_bf16_bf16_fp8x3 +// CHECK: algorithm = #mhlo.dot_algorithm +ENTRY %main.4 (Arg_0.1: bf16[2,2,2], Arg_1.2: bf16[2,2,2]) -> f32[2,2,2] { + %Arg_0.1 = bf16[2,2,2] parameter(0) + %Arg_1.2 = bf16[2,2,2] parameter(1) + ROOT %dot.3 = f32[2,2,2] dot(bf16[2,2,2] %Arg_0.1, bf16[2,2,2] %Arg_1.2), lhs_batch_dims={0}, lhs_contracting_dims={2}, rhs_batch_dims={0}, rhs_contracting_dims={1}, algorithm=dot_bf16_bf16_fp8x3 +} + +// ----- + +HloModule dot_algorithm_bf16_bf16_fp8x4, entry_computation_layout={(bf16[2,2,2]{2,1,0}, bf16[2,2,2]{2,1,0})->f32[2,2,2]{2,1,0}} + +// CHECK-LABEL: module @dot_algorithm_bf16_bf16_fp8x4 +// CHECK: algorithm = #mhlo.dot_algorithm +ENTRY %main.4 (Arg_0.1: bf16[2,2,2], Arg_1.2: bf16[2,2,2]) -> f32[2,2,2] { + %Arg_0.1 = bf16[2,2,2] parameter(0) + %Arg_1.2 = bf16[2,2,2] parameter(1) + ROOT %dot.3 = f32[2,2,2] dot(bf16[2,2,2] %Arg_0.1, bf16[2,2,2] %Arg_1.2), lhs_batch_dims={0}, lhs_contracting_dims={2}, rhs_batch_dims={0}, rhs_contracting_dims={1}, algorithm=dot_bf16_bf16_fp8x4 +} diff --git a/third_party/xla/xla/hlo/translate/mhlo_to_hlo/attribute_exporter.cc b/third_party/xla/xla/hlo/translate/mhlo_to_hlo/attribute_exporter.cc index 9ded7442e11c42..c49b88d3612efe 100644 --- a/third_party/xla/xla/hlo/translate/mhlo_to_hlo/attribute_exporter.cc +++ b/third_party/xla/xla/hlo/translate/mhlo_to_hlo/attribute_exporter.cc @@ -469,6 +469,10 @@ absl::StatusOr ConvertDotAlgorithm( return xla::PrecisionConfig::ALG_DOT_BF16_BF16_F32; case mlir::hlo::detail::KnownDotAlgorithm::BF16_BF16_F32_X3: return xla::PrecisionConfig::ALG_DOT_BF16_BF16_F32_X3; + case mlir::hlo::detail::KnownDotAlgorithm::F8E4M3FN_F8E4M3FN_F32_X3: + return xla::PrecisionConfig::ALG_DOT_BF16_BF16_FP8X3; + case mlir::hlo::detail::KnownDotAlgorithm::F8E4M3FN_F8E4M3FN_F32_X4: + return xla::PrecisionConfig::ALG_DOT_BF16_BF16_FP8X4; case mlir::hlo::detail::KnownDotAlgorithm::BF16_BF16_F32_X6: return xla::PrecisionConfig::ALG_DOT_BF16_BF16_F32_X6; case mlir::hlo::detail::KnownDotAlgorithm::BF16_BF16_F32_X9: @@ -511,6 +515,10 @@ absl::StatusOr ConvertDotAlgorithm( return xla::PrecisionConfig::ALG_DOT_BF16_BF16_F32; case mlir::hlo::detail::KnownDotAlgorithm::BF16_BF16_F32_X3: return xla::PrecisionConfig::ALG_DOT_BF16_BF16_F32_X3; + case mlir::hlo::detail::KnownDotAlgorithm::F8E4M3FN_F8E4M3FN_F32_X3: + return xla::PrecisionConfig::ALG_DOT_BF16_BF16_FP8X3; + case mlir::hlo::detail::KnownDotAlgorithm::F8E4M3FN_F8E4M3FN_F32_X4: + return xla::PrecisionConfig::ALG_DOT_BF16_BF16_FP8X4; case mlir::hlo::detail::KnownDotAlgorithm::BF16_BF16_F32_X6: return xla::PrecisionConfig::ALG_DOT_BF16_BF16_F32_X6; case mlir::hlo::detail::KnownDotAlgorithm::BF16_BF16_F32_X9: diff --git a/third_party/xla/xla/hlo/translate/mhlo_to_hlo/tests/attributes.mlir b/third_party/xla/xla/hlo/translate/mhlo_to_hlo/tests/attributes.mlir index d0c63e7d0ec6cd..3458d98d575349 100644 --- a/third_party/xla/xla/hlo/translate/mhlo_to_hlo/tests/attributes.mlir +++ b/third_party/xla/xla/hlo/translate/mhlo_to_hlo/tests/attributes.mlir @@ -299,3 +299,51 @@ module @dot_algorithm_f64_f64_f64 { }> : (tensor<2x2x2xf64>, tensor<2x2x2xf64>) -> tensor<2x2x2xf64> return %0 : tensor<2x2x2xf64> } } + +// ----- + +// CHECK-LABEL: HloModule dot_algorithm_bf16_bf16_fp8x3 +module @dot_algorithm_bf16_bf16_fp8x3 { + func.func @main(%arg0: tensor<2x2x2xbf16>, %arg1: tensor<2x2x2xbf16>) -> tensor<2x2x2xf32> { + // CHECK: %[[ARG0:.+]] = bf16[2,2,2] parameter(0) + // CHECK: %[[ARG1:.+]] = bf16[2,2,2] parameter(1) + // CHECK: f32[2,2,2] dot(%[[ARG0]], %[[ARG1]]), {{.*}}, algorithm=dot_bf16_bf16_fp8x3 + %0 = "mhlo.dot_general"(%arg0, %arg1) <{ + dot_dimension_numbers = #mhlo.dot, + precision_config = [#mhlo, #mhlo], + algorithm = #mhlo.dot_algorithm< + lhs_precision_type = f8E4M3FN, + rhs_precision_type = f8E4M3FN, + accumulation_type = f32, + lhs_component_count = 1, + rhs_component_count = 1, + num_primitive_operations = 3, + allow_imprecise_accumulation = false + > + }> : (tensor<2x2x2xbf16>, tensor<2x2x2xbf16>) -> tensor<2x2x2xf32> return %0 : tensor<2x2x2xf32> + } +} + +// ----- + +// CHECK-LABEL: HloModule dot_algorithm_bf16_bf16_fp8x4 +module @dot_algorithm_bf16_bf16_fp8x4 { + func.func @main(%arg0: tensor<2x2x2xbf16>, %arg1: tensor<2x2x2xbf16>) -> tensor<2x2x2xf32> { + // CHECK: %[[ARG0:.+]] = bf16[2,2,2] parameter(0) + // CHECK: %[[ARG1:.+]] = bf16[2,2,2] parameter(1) + // CHECK: f32[2,2,2] dot(%[[ARG0]], %[[ARG1]]), {{.*}}, algorithm=dot_bf16_bf16_fp8x4 + %0 = "mhlo.dot_general"(%arg0, %arg1) <{ + dot_dimension_numbers = #mhlo.dot, + precision_config = [#mhlo, #mhlo], + algorithm = #mhlo.dot_algorithm< + lhs_precision_type = f8E4M3FN, + rhs_precision_type = f8E4M3FN, + accumulation_type = f32, + lhs_component_count = 1, + rhs_component_count = 1, + num_primitive_operations = 4, + allow_imprecise_accumulation = false + > + }> : (tensor<2x2x2xbf16>, tensor<2x2x2xbf16>) -> tensor<2x2x2xf32> return %0 : tensor<2x2x2xf32> + } +} diff --git a/third_party/xla/xla/mlir_hlo/tests/Dialect/mhlo/hlo-legalize-to-stablehlo.mlir b/third_party/xla/xla/mlir_hlo/tests/Dialect/mhlo/hlo-legalize-to-stablehlo.mlir index 3237af6717d2fa..0fb9e69a7b5e83 100644 --- a/third_party/xla/xla/mlir_hlo/tests/Dialect/mhlo/hlo-legalize-to-stablehlo.mlir +++ b/third_party/xla/xla/mlir_hlo/tests/Dialect/mhlo/hlo-legalize-to-stablehlo.mlir @@ -1022,6 +1022,84 @@ func.func @op_dot_general_algorithm(%arg0: tensor<8x8x16xf32>, %arg1: tensor<8x1 func.return %0 : tensor<8x8x8xf32> } +// CHECK-LABEL: "op_dot_general_algorithm_fp8x3" +func.func @op_dot_general_algorithm_fp8x3(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { + // CHECK: "stablehlo.dot_general"([[ARG0:%arg[0-9]+]], [[ARG1:%arg[0-9]+]]) <{ + // CHECK-SAME: algorithm = #stablehlo.dot_algorithm< + // CHECK-SAME: lhs_precision_type = f8E4M3FN, + // CHECK-SAME: rhs_precision_type = f8E4M3FN, + // CHECK-SAME: accumulation_type = f32, + // CHECK-SAME: lhs_component_count = 1, + // CHECK-SAME: rhs_component_count = 1, + // CHECK-SAME: num_primitive_operations = 3, + // CHECK-SAME: allow_imprecise_accumulation = false + // CHECK-SAME: >, + // CHECK-SAME: dot_dimension_numbers = #stablehlo.dot< + // CHECK-SAME: lhs_batching_dimensions = [0], + // CHECK-SAME: rhs_batching_dimensions = [0], + // CHECK-SAME: lhs_contracting_dimensions = [2], + // CHECK-SAME: rhs_contracting_dimensions = [1] + // CHECK-SAME: > + // CHECK-SAME: }> : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> + %0 = "mhlo.dot_general"(%arg0, %arg1) { + dot_dimension_numbers = #mhlo.dot< + lhs_batching_dimensions = [0], + lhs_contracting_dimensions = [2], + rhs_batching_dimensions = [0], + rhs_contracting_dimensions = [1] + >, + algorithm = #mhlo.dot_algorithm< + lhs_precision_type = f8E4M3FN, + rhs_precision_type = f8E4M3FN, + accumulation_type = f32, + lhs_component_count = 1, + rhs_component_count = 1, + num_primitive_operations = 3, + allow_imprecise_accumulation = false + > + } : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> + func.return %0 : tensor<8x8x8xf32> +} + +// CHECK-LABEL: "op_dot_general_algorithm_fp8x4" +func.func @op_dot_general_algorithm_fp8x4(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { + // CHECK: "stablehlo.dot_general"([[ARG0:%arg[0-9]+]], [[ARG1:%arg[0-9]+]]) <{ + // CHECK-SAME: algorithm = #stablehlo.dot_algorithm< + // CHECK-SAME: lhs_precision_type = f8E4M3FN, + // CHECK-SAME: rhs_precision_type = f8E4M3FN, + // CHECK-SAME: accumulation_type = f32, + // CHECK-SAME: lhs_component_count = 1, + // CHECK-SAME: rhs_component_count = 1, + // CHECK-SAME: num_primitive_operations = 4, + // CHECK-SAME: allow_imprecise_accumulation = false + // CHECK-SAME: >, + // CHECK-SAME: dot_dimension_numbers = #stablehlo.dot< + // CHECK-SAME: lhs_batching_dimensions = [0], + // CHECK-SAME: rhs_batching_dimensions = [0], + // CHECK-SAME: lhs_contracting_dimensions = [2], + // CHECK-SAME: rhs_contracting_dimensions = [1] + // CHECK-SAME: > + // CHECK-SAME: }> : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> + %0 = "mhlo.dot_general"(%arg0, %arg1) { + dot_dimension_numbers = #mhlo.dot< + lhs_batching_dimensions = [0], + lhs_contracting_dimensions = [2], + rhs_batching_dimensions = [0], + rhs_contracting_dimensions = [1] + >, + algorithm = #mhlo.dot_algorithm< + lhs_precision_type = f8E4M3FN, + rhs_precision_type = f8E4M3FN, + accumulation_type = f32, + lhs_component_count = 1, + rhs_component_count = 1, + num_primitive_operations = 4, + allow_imprecise_accumulation = false + > + } : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> + func.return %0 : tensor<8x8x8xf32> +} + // CHECK-LABEL: "op_dot" func.func @op_dot(%arg0: tensor<8x16xf32>, %arg1: tensor<16x8xf32>) -> tensor<8x8xf32> { // CHECK: "stablehlo.dot"([[ARG0:%arg[0-9]+]], [[ARG1:%arg[0-9]+]]) <{ diff --git a/third_party/xla/xla/mlir_hlo/tests/Dialect/mhlo/stablehlo-legalize-to-hlo.mlir b/third_party/xla/xla/mlir_hlo/tests/Dialect/mhlo/stablehlo-legalize-to-hlo.mlir index 2cd7b908c21789..22e2552c0bbb49 100644 --- a/third_party/xla/xla/mlir_hlo/tests/Dialect/mhlo/stablehlo-legalize-to-hlo.mlir +++ b/third_party/xla/xla/mlir_hlo/tests/Dialect/mhlo/stablehlo-legalize-to-hlo.mlir @@ -1131,6 +1131,88 @@ func.func @op_dot_general_algorithm(%arg0: tensor<8x8x16xf32>, %arg1: tensor<8x1 // ----- +// CHECK-LABEL: "op_dot_general_algorithm_fp8x3" +func.func @op_dot_general_algorithm_fp8x3(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { + // CHECK: "mhlo.dot_general"([[ARG0:%arg[0-9]+]], [[ARG1:%arg[0-9]+]]) <{ + // CHECK-SAME: algorithm = #mhlo.dot_algorithm< + // CHECK-SAME: lhs_precision_type = f8E4M3FN, + // CHECK-SAME: rhs_precision_type = f8E4M3FN, + // CHECK-SAME: accumulation_type = f32, + // CHECK-SAME: lhs_component_count = 1, + // CHECK-SAME: rhs_component_count = 1, + // CHECK-SAME: num_primitive_operations = 3, + // CHECK-SAME: allow_imprecise_accumulation = false + // CHECK-SAME: >, + // CHECK-SAME: dot_dimension_numbers = #mhlo.dot< + // CHECK-SAME: lhs_batching_dimensions = [0], + // CHECK-SAME: rhs_batching_dimensions = [0], + // CHECK-SAME: lhs_contracting_dimensions = [2], + // CHECK-SAME: rhs_contracting_dimensions = [1] + // CHECK-SAME: > + // CHECK-SAME: }> : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> + %0 = "stablehlo.dot_general"(%arg0, %arg1) { + dot_dimension_numbers = #stablehlo.dot< + lhs_batching_dimensions = [0], + lhs_contracting_dimensions = [2], + rhs_batching_dimensions = [0], + rhs_contracting_dimensions = [1] + >, + algorithm = #stablehlo.dot_algorithm< + lhs_precision_type = f8E4M3FN, + rhs_precision_type = f8E4M3FN, + accumulation_type = f32, + lhs_component_count = 1, + rhs_component_count = 1, + num_primitive_operations = 3, + allow_imprecise_accumulation = false + > + } : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> + func.return %0 : tensor<8x8x8xf32> +} + +// ----- + +// CHECK-LABEL: "op_dot_general_algorithm_fp8x4" +func.func @op_dot_general_algorithm_fp8x4(%arg0: tensor<8x8x16xbf16>, %arg1: tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> { + // CHECK: "mhlo.dot_general"([[ARG0:%arg[0-9]+]], [[ARG1:%arg[0-9]+]]) <{ + // CHECK-SAME: algorithm = #mhlo.dot_algorithm< + // CHECK-SAME: lhs_precision_type = f8E4M3FN, + // CHECK-SAME: rhs_precision_type = f8E4M3FN, + // CHECK-SAME: accumulation_type = f32, + // CHECK-SAME: lhs_component_count = 1, + // CHECK-SAME: rhs_component_count = 1, + // CHECK-SAME: num_primitive_operations = 4, + // CHECK-SAME: allow_imprecise_accumulation = false + // CHECK-SAME: >, + // CHECK-SAME: dot_dimension_numbers = #mhlo.dot< + // CHECK-SAME: lhs_batching_dimensions = [0], + // CHECK-SAME: rhs_batching_dimensions = [0], + // CHECK-SAME: lhs_contracting_dimensions = [2], + // CHECK-SAME: rhs_contracting_dimensions = [1] + // CHECK-SAME: > + // CHECK-SAME: }> : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> + %0 = "stablehlo.dot_general"(%arg0, %arg1) { + dot_dimension_numbers = #stablehlo.dot< + lhs_batching_dimensions = [0], + lhs_contracting_dimensions = [2], + rhs_batching_dimensions = [0], + rhs_contracting_dimensions = [1] + >, + algorithm = #stablehlo.dot_algorithm< + lhs_precision_type = f8E4M3FN, + rhs_precision_type = f8E4M3FN, + accumulation_type = f32, + lhs_component_count = 1, + rhs_component_count = 1, + num_primitive_operations = 4, + allow_imprecise_accumulation = false + > + } : (tensor<8x8x16xbf16>, tensor<8x16x8xbf16>) -> tensor<8x8x8xf32> + func.return %0 : tensor<8x8x8xf32> +} + +// ----- + // CHECK-LABEL: "op_dot" func.func @op_dot(%arg0: tensor<8x16xf32>, %arg1: tensor<16x8xf32>) -> tensor<8x8xf32> { // CHECK: "mhlo.dot"([[ARG0:%arg[0-9]+]], [[ARG1:%arg[0-9]+]]) <{ diff --git a/third_party/xla/xla/service/algorithm_util.cc b/third_party/xla/xla/service/algorithm_util.cc index 97eecf661f4b4a..80b4200b1dec24 100644 --- a/third_party/xla/xla/service/algorithm_util.cc +++ b/third_party/xla/xla/service/algorithm_util.cc @@ -58,6 +58,8 @@ absl::StatusOr GetBlasComputationType( case PrecisionConfig::ALG_DOT_BF16_BF16_F32_X3: case PrecisionConfig::ALG_DOT_BF16_BF16_F32_X6: case PrecisionConfig::ALG_DOT_BF16_BF16_F32_X9: + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X3: + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X4: case PrecisionConfig::ALG_DOT_F32_F32_F32: case PrecisionConfig::ALG_DOT_TF32_TF32_F32_X3: @@ -88,6 +90,9 @@ absl::StatusOr> GetAllowedOperandsTypeForAlgorithm( case PrecisionConfig::ALG_DOT_BF16_BF16_BF16: case PrecisionConfig::ALG_DOT_BF16_BF16_F32: return std::vector{BF16}; + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X3: + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X4: + return std::vector{BF16, F32}; case PrecisionConfig::ALG_DOT_BF16_BF16_F32_X3: case PrecisionConfig::ALG_DOT_BF16_BF16_F32_X6: case PrecisionConfig::ALG_DOT_BF16_BF16_F32_X9: @@ -129,6 +134,8 @@ absl::StatusOr GetDotAccumulatorType( case PrecisionConfig::ALG_DOT_BF16_BF16_F32_X3: case PrecisionConfig::ALG_DOT_BF16_BF16_F32_X6: case PrecisionConfig::ALG_DOT_BF16_BF16_F32_X9: + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X3: + case PrecisionConfig::ALG_DOT_BF16_BF16_FP8X4: case PrecisionConfig::ALG_DOT_TF32_TF32_F32: case PrecisionConfig::ALG_DOT_TF32_TF32_F32_X3: case PrecisionConfig::ALG_DOT_F32_F32_F32: diff --git a/third_party/xla/xla/xla_data.proto b/third_party/xla/xla/xla_data.proto index 9cd793c6e2a362..604b7fa3f121fb 100644 --- a/third_party/xla/xla/xla_data.proto +++ b/third_party/xla/xla/xla_data.proto @@ -1434,8 +1434,14 @@ message PrecisionConfig { ALG_DOT_F32_F32_F32 = 11; ALG_DOT_F64_F64_F64 = 12; ALG_DOT_BF16_BF16_F32_X9 = 13; - - // Next: 14 + // An algorithm which uses 3 FP8 matmuls to emulate a BF16_BF16 dot product + // with F32 accumulation. + ALG_DOT_BF16_BF16_FP8X3 = 14; + // An algorithm which uses 4 FP8 matmuls to emulate a BF16_BF16 dot product + // with F32 accumulation. + ALG_DOT_BF16_BF16_FP8X4 = 15; + + // Next: 16 } repeated Precision operand_precision = 1; From e324dfeed71a7c57baee29b078946d7516c1322f Mon Sep 17 00:00:00 2001 From: Deqiang Chen Date: Tue, 29 Sep 2026 17:00:09 -0700 Subject: [PATCH 09/21] Use HeapSimulator to decide schedule-by-structure reruns and restore input schedule when we change schduler config. PiperOrigin-RevId: 990614610 --- third_party/xla/xla/service/BUILD | 1 + .../xla/service/latency_hiding_scheduler.cc | 181 +++++++++++++----- .../xla/service/latency_hiding_scheduler.h | 12 +- 3 files changed, 150 insertions(+), 44 deletions(-) diff --git a/third_party/xla/xla/service/BUILD b/third_party/xla/xla/service/BUILD index 81299291f0dcc5..7ed4d95dd64d5c 100644 --- a/third_party/xla/xla/service/BUILD +++ b/third_party/xla/xla/service/BUILD @@ -1267,6 +1267,7 @@ cc_library( "@com_google_absl//absl/strings:str_format", "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/synchronization", + "@com_google_absl//absl/time", "@com_google_absl//absl/types:span", "@com_googlesource_code_re2//:re2", ], diff --git a/third_party/xla/xla/service/latency_hiding_scheduler.cc b/third_party/xla/xla/service/latency_hiding_scheduler.cc index 3ec3029e575370..e819190eeab6da 100644 --- a/third_party/xla/xla/service/latency_hiding_scheduler.cc +++ b/third_party/xla/xla/service/latency_hiding_scheduler.cc @@ -46,6 +46,8 @@ limitations under the License. #include "absl/strings/str_join.h" #include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" +#include "absl/time/clock.h" +#include "absl/time/time.h" #include "absl/types/span.h" #include "re2/re2.h" #include "xla/hlo/analysis/alias_info.h" @@ -299,30 +301,79 @@ GetNumResourcesNeededForAnnotationWithKeepOriginalOrderAttrs( return max_resources_needed; } -int64_t EstimateFragmentationSize(HloModule* module, - const HloAliasAnalysis& alias_analysis, - const AliasInfo* alias_info) { - // Run heap simulator on the whole module to estimate the fragmentation size. - auto algorithm = std::make_unique>( - /*alignment=*/1); - BufferValue::SizeFunction size_fn = [](const BufferValue& buffer) -> int64_t { - const Shape& shape = buffer.shape(); - if (!shape.IsArray()) { - return 0; +struct HbmUsageEstimate { + // Peak heap size of the simulated allocation, including fragmentation. + int64_t heap_size = 0; + // heap_size minus the heap size of an ideal, fragmentation-free allocation. + int64_t fragmentation_size = 0; +}; + +// Runs a whole-module HeapSimulator on the current schedule. Only +// default-memory-space arrays are counted. +// +// If `fragmentation_only`, this is the legacy fragmentation estimate: all +// values are counted with their unpadded ShapeUtil::ByteSizeOf size and spatial +// packing; `shape_size_bytes` is ignored. +// +// Otherwise, this estimates the non-parameter memory, which is what memory +// limits typically apply to: HloBuffers containing an entry computation +// parameter get size 0 (HeapSimulator requires every value it allocates to be +// assigned, so they are zero-sized instead of being filtered out), sizes come +// from `shape_size_bytes`, and fast-merge packing is used, as in TPU buffer +// assignment. +absl::StatusOr EstimateHbmUsage( + const HloModule& module, const HloAliasAnalysis& alias_analysis, + const AliasInfo* alias_info, bool fragmentation_only, + const HloCostAnalysis::ShapeSizeFunction& shape_size_bytes = nullptr) { + using Heap = GlobalDecreasingSizeBestFitHeap; + absl::flat_hash_set excluded_value_ids; + if (!fragmentation_only) { + CHECK(shape_size_bytes != nullptr); + const HloComputation* entry = module.entry_computation(); + for (const HloBuffer& buffer : alias_analysis.buffers()) { + if (absl::c_any_of(buffer.values(), [&](const HloValue* value) { + const HloInstruction* def = value->defining_instruction(); + return def->opcode() == HloOpcode::kParameter && + def->parent() == entry; + })) { + for (const HloValue* value : buffer.values()) { + excluded_value_ids.insert(value->id()); + } + } } - if (!shape.has_layout()) { + } + BufferValue::SizeFunction size_fn = + [&](const BufferValue& buffer) -> int64_t { + if (excluded_value_ids.contains(buffer.id())) { return 0; } - if (shape.layout().memory_space() != Layout::kDefaultMemorySpace) { + const Shape& shape = buffer.shape(); + if (!shape.IsArray() || !shape.has_layout() || + shape.layout().memory_space() != Layout::kDefaultMemorySpace) { return 0; } - return ShapeUtil::ByteSizeOf(shape); + return fragmentation_only ? ShapeUtil::ByteSizeOf(shape) + : shape_size_bytes(shape); }; - auto result = - HeapSimulator::Run(std::move(algorithm), *module, module->schedule(), - alias_analysis, alias_info, &size_fn); - CHECK_OK(result.status()); - int64_t fragmentation_size = result.value().fragmentation_size; + ABSL_ASSIGN_OR_RETURN( + HeapSimulator::Result result, + HeapSimulator::Run( + std::make_unique(/*alignment=*/1, fragmentation_only + ? Heap::kSpatial + : Heap::kFastMerge), + module, module.schedule(), alias_analysis, alias_info, &size_fn)); + return HbmUsageEstimate{result.heap_size, result.fragmentation_size}; +} + +int64_t EstimateFragmentationSize(HloModule* module, + const HloAliasAnalysis& alias_analysis, + const AliasInfo* alias_info) { + // Run heap simulator on the whole module to estimate the fragmentation size. + absl::StatusOr estimate = + EstimateHbmUsage(*module, alias_analysis, alias_info, + /*fragmentation_only=*/true); + CHECK_OK(estimate.status()); + int64_t fragmentation_size = estimate->fragmentation_size; VLOG(3) << module->name() << ": Heap simulator estimated fragmentation size: " << fragmentation_size; return fragmentation_size > 0 ? fragmentation_size : 0; @@ -4446,6 +4497,21 @@ absl::StatusOr LatencyHidingScheduler::RunImpl( .ToString()); } } + AsyncTracker* async_tracker = scheduling_context_->GetAsyncTracker().get(); + const bool use_heap_simulator = + async_tracker->GetConfig().enable_schedule_by_structure; + // With schedule-by-structure, the first rerun starts from the same input + // schedule as the first try, so that it only differs in the scheduler + // configuration. + std::vector> + input_sequences; + if (use_heap_simulator && scheduler_core_->GetRerunTimes() > 0) { + input_sequences.reserve(computations_to_schedule_.size()); + for (HloComputation* computation : computations_to_schedule_) { + input_sequences.emplace_back(computation, + module->schedule().sequence(computation)); + } + } for (HloComputation* computation : computations_to_schedule_) { ABSL_ASSIGN_OR_RETURN(std::vector new_schedule, scheduler_core_->ScheduleComputation(computation)); @@ -4459,26 +4525,63 @@ absl::StatusOr LatencyHidingScheduler::RunImpl( scheduler_core_->GetSchedulingState().get()); scheduling_context_->GetAsyncTracker()->InvalidateCache(computation); } - int64_t fragmentation_size = - scheduling_context_->GetAsyncTracker() - ->GetConfig() - .estimate_fragmentation_size - ? EstimateFragmentationSize(module, - *scheduling_context_->GetAliasAnalysis(), - scheduling_context_->GetAliasInfo()) - : 0; uint64_t initial_memory_limit = scheduler_core_->GetMemoryLimit(); - for (int64_t iter = 0; iter < scheduler_core_->GetRerunTimes() && - scheduler_core_->GetMemoryPeak() + fragmentation_size > - initial_memory_limit; - iter++) { - LOG(INFO) << "LatencyHidingScheduler current memory usage: " - << scheduler_core_->GetMemoryPeak() + fragmentation_size - << " bytes, does not fit in initial limit: " + // Returns the memory usage of the current schedule that decides whether to + // rerun with a tighter memory limit. With schedule-by-structure, the LHS + // memory peak is not a reliable measure of the final memory usage, so a + // HeapSimulator estimate of the non-parameter memory (the part the memory + // limit applies to) is used instead. + auto estimate_memory_usage = [&]() -> absl::StatusOr { + if (!use_heap_simulator) { + int64_t fragmentation_size = + async_tracker->GetConfig().estimate_fragmentation_size + ? EstimateFragmentationSize( + module, *scheduling_context_->GetAliasAnalysis(), + scheduling_context_->GetAliasInfo()) + : 0; + return scheduler_core_->GetMemoryPeak() + fragmentation_size; + } + const absl::Time start = absl::Now(); + ABSL_ASSIGN_OR_RETURN( + HbmUsageEstimate estimate, + EstimateHbmUsage(*module, *scheduling_context_->GetAliasAnalysis(), + scheduling_context_->GetAliasInfo(), + /*fragmentation_only=*/false, + scheduling_context_->GetShapeSizeBytes())); + LOG(INFO) << "[" << name() + << "] LatencyHidingScheduler HeapSimulator memory estimate: " + << estimate.heap_size + << " bytes (fragmentation: " << estimate.fragmentation_size + << "). LHS memory peak: " << scheduler_core_->GetMemoryPeak() + << ". Estimate took " + << absl::FormatDuration(absl::Now() - start); + return estimate.heap_size; + }; + for (int64_t iter = 0; iter < scheduler_core_->GetRerunTimes(); iter++) { + ABSL_ASSIGN_OR_RETURN(int64_t memory_usage, estimate_memory_usage()); + if (static_cast(memory_usage) <= initial_memory_limit) { + break; + } + uint64_t new_limit = + static_cast(scheduler_core_->GetMemoryLimit() * 0.9); + if (use_heap_simulator) { + async_tracker->SetEnableCpdForSyncCollective(false); + if (iter == 0) { + // Restart from the original schedule when we changed the scheduler + // config. + for (const auto& [computation, sequence] : input_sequences) { + module->schedule().set_sequence(computation, sequence); + } + async_tracker->InvalidateCache(); + } + } + LOG(INFO) << "[" << name() + << "] LatencyHidingScheduler current memory usage: " + << memory_usage << " bytes, does not fit in initial limit: " << initial_memory_limit << ". Setting the new limit to " - << static_cast(scheduler_core_->GetMemoryLimit() * 0.9); + << new_limit; ABSL_RETURN_IF_ERROR(scheduler_core_->InitializeScheduler(module)); - scheduler_core_->SetMemoryLimit(scheduler_core_->GetMemoryLimit() * 0.9); + scheduler_core_->SetMemoryLimit(new_limit); for (HloComputation* computation : computations_to_schedule_) { ABSL_ASSIGN_OR_RETURN(std::vector new_schedule, scheduler_core_->ScheduleComputation(computation)); @@ -4490,14 +4593,6 @@ absl::StatusOr LatencyHidingScheduler::RunImpl( scheduler_core_->GetSchedulingState().get()); scheduling_context_->GetAsyncTracker()->InvalidateCache(computation); } - fragmentation_size = - scheduling_context_->GetAsyncTracker() - ->GetConfig() - .estimate_fragmentation_size - ? EstimateFragmentationSize( - module, *scheduling_context_->GetAliasAnalysis(), - scheduling_context_->GetAliasInfo()) - : 0; } LOG(INFO) << "[" << name() << "]" << " LatencyHidingScheduler current memory usage: " diff --git a/third_party/xla/xla/service/latency_hiding_scheduler.h b/third_party/xla/xla/service/latency_hiding_scheduler.h index 9b3b39b11fd592..46bdc417692ff1 100644 --- a/third_party/xla/xla/service/latency_hiding_scheduler.h +++ b/third_party/xla/xla/service/latency_hiding_scheduler.h @@ -208,6 +208,10 @@ struct SchedulerConfig { bool top_down_scheduling = false; // If true, enable schedule by structure. bool enable_schedule_by_structure = false; + // If true (and enable_schedule_by_structure is true), synchronous + // collectives are used as schedule anchors by the target-specific + // schedule-by-structure (critical path depth) heuristic. + bool enable_cpd_for_sync_collective = true; // If set, only log computations that match the given regular expression. std::string log_computation_re; }; @@ -474,6 +478,12 @@ class AsyncTracker { const SchedulerConfig& GetConfig() const { return config_; } + // Overrides SchedulerConfig::enable_cpd_for_sync_collective. Used by the + // scheduler to reschedule with a different schedule-by-structure setting. + void SetEnableCpdForSyncCollective(bool enable) { + config_.enable_cpd_for_sync_collective = enable; + } + // Clears the cache of per-computation resource maps. This is needed when, // e.g., we modify the schedule of a computation, which could change the // resource usage of the computation. @@ -521,7 +531,7 @@ class AsyncTracker { GetCanonicalAsyncOpFunc get_canonical_async_op_; protected: - const SchedulerConfig config_; + SchedulerConfig config_; mutable absl::Mutex resources_cache_mu_; mutable absl::flat_hash_map> From e8fecb237ba91b8bfbbb95ba10d0aceaca9d82bd Mon Sep 17 00:00:00 2001 From: Emilio Cota Date: Tue, 29 Sep 2026 17:26:55 -0700 Subject: [PATCH 10/21] [xla:cpu] sink same-block, contract, single-use fmul before fadd/fsub So that sanitizers won't preclude LLVM passes from combining them into an fma. Sanitizer instrumentation can break basic blocks, which prevents LLVM from fusing FMA's since FMA's must come from same-basic-block pairs. This change is necessary to maintain the same numerics when we have full msan support, which is upcoming. PiperOrigin-RevId: 990627586 --- .../xla/backends/cpu/codegen/ir_compiler.cc | 4 + third_party/xla/xla/service/llvm_ir/BUILD | 12 + .../xla/xla/service/llvm_ir/llvm_util.cc | 25 ++ .../xla/xla/service/llvm_ir/llvm_util.h | 16 + .../xla/xla/service/llvm_ir/llvm_util_test.cc | 274 ++++++++++++++++++ 5 files changed, 331 insertions(+) create mode 100644 third_party/xla/xla/service/llvm_ir/llvm_util_test.cc diff --git a/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc index 05e845e75ef9f4..c59d2cded0e085 100644 --- a/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc +++ b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc @@ -473,6 +473,10 @@ llvm::Error IrCompiler::RunIrPasses(llvm::Module& module, // Must run after all optimization passes: middle-end passes behave // differently on instructions that already carry `contract`. llvm_ir::SetAllowContractOnFpArithmetic(module); + // Must run after `contract` is set and before sanitizer instrumentation, so + // that instrumentation cannot split contractable fmul/fadd pairs into + // separate basic blocks. + llvm_ir::SinkContractableFMulToFAddFSub(module); // Sanitizer instrumentation must be the last IR transformation. if (options_.dfsan_enabled) { diff --git a/third_party/xla/xla/service/llvm_ir/BUILD b/third_party/xla/xla/service/llvm_ir/BUILD index 59efa7e58e685f..7cc519dd70ff4a 100644 --- a/third_party/xla/xla/service/llvm_ir/BUILD +++ b/third_party/xla/xla/service/llvm_ir/BUILD @@ -101,6 +101,18 @@ cc_library( ], ) +xla_cc_test( + name = "llvm_util_test", + srcs = ["llvm_util_test.cc"], + deps = [ + ":llvm_util", + "@com_google_googletest//:gtest_main", + "@llvm-project//llvm:AsmParser", + "@llvm-project//llvm:Support", + "@llvm-project//llvm:ir_headers", + ], +) + cc_library( name = "llvm_type_conversion_util", hdrs = ["llvm_type_conversion_util.h"], diff --git a/third_party/xla/xla/service/llvm_ir/llvm_util.cc b/third_party/xla/xla/service/llvm_ir/llvm_util.cc index 336ef100154ec6..7a2428d2ebcbfe 100644 --- a/third_party/xla/xla/service/llvm_ir/llvm_util.cc +++ b/third_party/xla/xla/service/llvm_ir/llvm_util.cc @@ -704,6 +704,31 @@ void SetAllowContractOnFpArithmetic(llvm::Module& module) { } } +void SinkContractableFMulToFAddFSub(llvm::Module& module) { + for (llvm::Function& function : module) { + for (llvm::Instruction& instruction : llvm::instructions(function)) { + const unsigned opcode = instruction.getOpcode(); + if ((opcode != llvm::Instruction::FAdd && + opcode != llvm::Instruction::FSub) || + !instruction.hasAllowContract()) { + continue; + } + for (llvm::Value* operand : instruction.operands()) { + auto* fmul = llvm::dyn_cast(operand); + // `fmul` is an operand of `instruction` and has exactly one use, so it + // dominates and (being in the same block) precedes `instruction`. + // Moving it forward therefore never disturbs the iteration above. + if (fmul != nullptr && fmul->getOpcode() == llvm::Instruction::FMul && + fmul->hasAllowContract() && fmul->hasOneUse() && + fmul->getParent() == instruction.getParent() && + fmul->getNextNode() != &instruction) { + fmul->moveBefore(instruction.getIterator()); + } + } + } + } +} + std::map MergeMetadata( llvm::LLVMContext* context, const std::map& a, const std::map& b) { diff --git a/third_party/xla/xla/service/llvm_ir/llvm_util.h b/third_party/xla/xla/service/llvm_ir/llvm_util.h index bd0c7349b34f55..639901ac5be90a 100644 --- a/third_party/xla/xla/service/llvm_ir/llvm_util.h +++ b/third_party/xla/xla/service/llvm_ir/llvm_util.h @@ -332,6 +332,22 @@ llvm::FastMathFlags GetCpuFastMathFlags(const HloModuleConfig& module_config); // surrounding expression tree (Reassociate.cpp:191). void SetAllowContractOnFpArithmetic(llvm::Module& module); +// Moves every single-use `fmul contract` to immediately before its +// `fadd contract` / `fsub contract` user when both are in the same basic block. +// +// Sanitizer instrumentation (e.g. msan's checks on loads/stores) inserts +// branches that split blocks. If this instrumentation is inserted between +// the fmul and its user, they no longer are in the same basic block, which +// prevents LLVM from fusing them into an fma. This transformation avoids that. +// +// Without sanitizers this changes the order of pure arithmetic within a block. +// The only observable effect is a different (order-based) tie-breaking in +// instruction scheduling. +// +// This is a workaround until XLA:CPU stops relying on LLVM passes to emit +// fma's; see b/560320144. +void SinkContractableFMulToFAddFSub(llvm::Module& module); + // Computes a conservative union of the metadata in "a" and "b". For // aliasing-related metadata, this means the result can be applied to // instructions whose aliasing relationship can be described either by "a" *or* diff --git a/third_party/xla/xla/service/llvm_ir/llvm_util_test.cc b/third_party/xla/xla/service/llvm_ir/llvm_util_test.cc new file mode 100644 index 00000000000000..8d8b0325fa4efd --- /dev/null +++ b/third_party/xla/xla/service/llvm_ir/llvm_util_test.cc @@ -0,0 +1,274 @@ +/* Copyright 2026 The OpenXLA Authors. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#include "xla/service/llvm_ir/llvm_util.h" + +#include +#include +#include + +#include +#include "llvm/ADT/StringRef.h" +#include "llvm/AsmParser/Parser.h" +#include "llvm/IR/Function.h" +#include "llvm/IR/InstIterator.h" +#include "llvm/IR/Instruction.h" +#include "llvm/IR/LLVMContext.h" +#include "llvm/IR/Module.h" +#include "llvm/Support/SourceMgr.h" + +namespace xla::llvm_ir { +namespace { + +llvm::Instruction* FindInstruction(llvm::Function& fn, llvm::StringRef name) { + for (llvm::Instruction& inst : llvm::instructions(fn)) { + if (inst.getName() == name) { + return &inst; + } + } + return nullptr; +} + +std::vector InstructionOrder( + const llvm::Function& fn) { + std::vector order; + for (const llvm::Instruction& inst : llvm::instructions(fn)) { + order.push_back(&inst); + } + return order; +} + +struct SinkFMulTestCase { + std::string test_name; + std::string ir; + bool should_sink; +}; + +class SinkContractableFMulToFAddFSubTest + : public ::testing::TestWithParam {}; + +TEST_P(SinkContractableFMulToFAddFSubTest, + SinksOnlyContractableSingleUsePairs) { + const SinkFMulTestCase& tc = GetParam(); + llvm::LLVMContext context; + llvm::SMDiagnostic diagnostic; + std::unique_ptr module = + llvm::parseAssemblyString(tc.ir, diagnostic, context); + ASSERT_NE(module, nullptr) << diagnostic.getMessage().str(); + + llvm::Function* fn = module->getFunction("test_fn"); + ASSERT_NE(fn, nullptr); + const std::vector order_before = + InstructionOrder(*fn); + + SinkContractableFMulToFAddFSub(*module); + + llvm::Instruction* mul = FindInstruction(*fn, "mul"); + llvm::Instruction* user = FindInstruction(*fn, "user"); + ASSERT_NE(mul, nullptr); + ASSERT_NE(user, nullptr); + if (tc.should_sink) { + EXPECT_EQ(mul->getNextNode(), user) << DumpToString(fn); + if (llvm::Instruction* mul0 = FindInstruction(*fn, "mul0")) { + EXPECT_EQ(mul0->getNextNode(), mul) << DumpToString(fn); + } + } else { + // Nothing at all should have moved. + EXPECT_EQ(InstructionOrder(*fn), order_before) << DumpToString(fn); + } +} + +INSTANTIATE_TEST_SUITE_P( + SinkContractableFMulToFAddFSubTests, SinkContractableFMulToFAddFSubTest, + ::testing::Values(SinkFMulTestCase{ + /*test_name=*/"FAddOperand0SingleUse", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %mul = fmul contract float %a, %b + %load = load float, ptr %p, align 4 + %user = fadd contract float %mul, %c + ret float %user + } + )", + /*should_sink=*/true, + }, + SinkFMulTestCase{ + /*test_name=*/"FAddOperand1SingleUse", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %mul = fmul contract float %a, %b + %load = load float, ptr %p, align 4 + %user = fadd contract float %c, %mul + ret float %user + } + )", + /*should_sink=*/true, + }, + SinkFMulTestCase{ + /*test_name=*/"FSubOperand0SingleUse", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %mul = fmul contract float %a, %b + %load = load float, ptr %p, align 4 + %user = fsub contract float %mul, %c + ret float %user + } + )", + /*should_sink=*/true, + }, + SinkFMulTestCase{ + /*test_name=*/"FSubOperand1SingleUse", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %mul = fmul contract float %a, %b + %load = load float, ptr %p, align 4 + %user = fsub contract float %c, %mul + ret float %user + } + )", + /*should_sink=*/true, + }, + SinkFMulTestCase{ + /*test_name=*/"FAddBothOperandsSingleUse", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, float %d, ptr %p) { + entry: + %mul0 = fmul contract float %a, %b + %mul = fmul contract float %c, %d + %load = load float, ptr %p, align 4 + %user = fadd contract float %mul0, %mul + ret float %user + } + )", + /*should_sink=*/true, + }, + SinkFMulTestCase{ + /*test_name=*/"VectorSingleUse", + /*ir=*/R"( + define <4 x float> @test_fn(<4 x float> %a, <4 x float> %b, <4 x float> %c, ptr %p) { + entry: + %mul = fmul contract <4 x float> %a, %b + %load = load float, ptr %p, align 4 + %user = fadd contract <4 x float> %mul, %c + ret <4 x float> %user + } + )", + /*should_sink=*/true, + }, + SinkFMulTestCase{ + /*test_name=*/"AlreadyAdjacentUnchanged", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %load = load float, ptr %p, align 4 + %mul = fmul contract float %a, %b + %user = fadd contract float %mul, %c + ret float %user + } + )", + /*should_sink=*/true, + }, + SinkFMulTestCase{ + /*test_name=*/"NoContractOnFMulNotSunk", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %mul = fmul float %a, %b + %load = load float, ptr %p, align 4 + %user = fadd contract float %mul, %c + ret float %user + } + )", + /*should_sink=*/false, + }, + SinkFMulTestCase{ + /*test_name=*/"NoContractOnUserNotSunk", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %mul = fmul contract float %a, %b + %load = load float, ptr %p, align 4 + %user = fadd float %mul, %c + ret float %user + } + )", + /*should_sink=*/false, + }, + SinkFMulTestCase{ + /*test_name=*/"MultiUseNotSunk", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %mul = fmul contract float %a, %b + %load = load float, ptr %p, align 4 + %user = fadd contract float %mul, %c + %extra = fadd contract float %mul, %load + ret float %extra + } + )", + /*should_sink=*/false, + }, + SinkFMulTestCase{ + /*test_name=*/"FDivUserNotSunk", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %mul = fmul contract float %a, %b + %load = load float, ptr %p, align 4 + %user = fdiv contract float %mul, %c + ret float %user + } + )", + /*should_sink=*/false, + }, + SinkFMulTestCase{ + /*test_name=*/"FMulUserNotSunk", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, ptr %p) { + entry: + %mul = fmul contract float %a, %b + %load = load float, ptr %p, align 4 + %user = fmul contract float %mul, %c + ret float %user + } + )", + /*should_sink=*/false, + }, + SinkFMulTestCase{ + /*test_name=*/"DifferentBasicBlocksNotSunk", + /*ir=*/R"( + define float @test_fn(float %a, float %b, float %c, i1 %cond) { + entry: + %mul = fmul contract float %a, %b + br i1 %cond, label %then, label %else + then: + %user = fadd contract float %mul, %c + ret float %user + else: + ret float %c + } + )", + /*should_sink=*/false, + }), + [](const ::testing::TestParamInfo& info) { + return info.param.test_name; + }); + +} // namespace +} // namespace xla::llvm_ir From fdf964bed383aa49c68eb71c3c65c4b4c5bf8984 Mon Sep 17 00:00:00 2001 From: Pedro Goncalves Mokarzel Date: Tue, 29 Sep 2026 18:01:42 -0700 Subject: [PATCH 11/21] Previously, `buildMaxAndArgmaxBody` computed `selected_value` using `SelectOp` with the index tie-breaking condition rather than `MaxOp`. While `SelectOp` produces the same value for normal numbers, it has different IEEE-754 semantics for `NaN`s (dropping `NaN` when it appears on the LHS) and signed zeros, which caused divergent `NaN` propagation behavior and prevented XLA from simplifying unused-index reductions to `kMaximum`. Rely on `MaxOp` for `selected_value` while keeping `SelectOp` for `selected_index`, and add a unit test in `StablehloBuilderTest` verifying the generated reducer body. PiperOrigin-RevId: 990642043 --- .../xla/third_party/stablehlo/temporary.patch | 198 +++++++++++++++++- 1 file changed, 197 insertions(+), 1 deletion(-) diff --git a/third_party/xla/third_party/stablehlo/temporary.patch b/third_party/xla/third_party/stablehlo/temporary.patch index d452af844efe30..6a05d2e1510d36 100644 --- a/third_party/xla/third_party/stablehlo/temporary.patch +++ b/third_party/xla/third_party/stablehlo/temporary.patch @@ -856,6 +856,17 @@ diff --ruN a/stablehlo/stablehlo/dialect/ChloOps.cpp b/stablehlo/stablehlo/diale } if (!hlo::isCompatibleForHloTypeInference( argType, bodyBlock.getArgument(i).getType())) { +diff --ruN a/stablehlo/stablehlo/dialect/ChloOps.td b/stablehlo/stablehlo/dialect/ChloOps.td +--- stablehlo/stablehlo/dialect/ChloOps.td ++++ stablehlo/stablehlo/dialect/ChloOps.td +@@ -54,6 +54,7 @@ + and provide conversion patterns to fully materialize into lower level + dialects. + }]; ++ let useStrictPropertiesInAssemblyFormat = 0; + } + + class CHLO_Op traits> : diff --ruN a/stablehlo/stablehlo/dialect/Register.cpp b/stablehlo/stablehlo/dialect/Register.cpp --- stablehlo/stablehlo/dialect/Register.cpp +++ stablehlo/stablehlo/dialect/Register.cpp @@ -971,10 +982,71 @@ diff --ruN a/stablehlo/stablehlo/dialect/StablehloEnums.td b/stablehlo/stablehlo +} #endif // STABLEHLO_DIALECT_STABLEHLO_ENUMS +diff --ruN a/stablehlo/stablehlo/dialect/StablehloOps.cpp b/stablehlo/stablehlo/dialect/StablehloOps.cpp +--- stablehlo/stablehlo/dialect/StablehloOps.cpp ++++ stablehlo/stablehlo/dialect/StablehloOps.cpp +@@ -4332,6 +4332,23 @@ + ReturnOp::create(*builder, loc, compare); + } + ++// For floating-point types, returns a predicate that selects lhs when lhs is ++// NaN and either rhs is not NaN or lhs_index < rhs_index (lt_index_pred), so ++// that index selection matches MaxOp/MinOp NaN propagation and tie-breaking. ++static Value buildNaNLhsSelectionCondition(OpBuilder& builder, Location loc, ++ Value lhs_value, Value rhs_value, ++ Value lt_index_pred) { ++ auto lhs_is_nan = CompareOp::create(builder, loc, lhs_value, lhs_value, ++ ComparisonDirection::NE) ++ .getResult(); ++ auto rhs_not_nan = CompareOp::create(builder, loc, rhs_value, rhs_value, ++ ComparisonDirection::EQ) ++ .getResult(); ++ auto nan_lhs_win = ++ OrOp::create(builder, loc, rhs_not_nan, lt_index_pred).getResult(); ++ return AndOp::create(builder, loc, lhs_is_nan, nan_lhs_win).getResult(); ++} ++ + void buildMaxAndArgmaxBody(Type elementType, Type indices_type, Region& body, + OpBuilder& builder) { + OpBuilder::InsertionGuard guard(builder); +@@ -4366,15 +4383,18 @@ + // Final lhs Selection Condition: (gt_pred) OR (tie_breaker_condition) + auto final_lhs_condition = + OrOp::create(builder, loc, gt_pred, tie_breaker_condition).getResult(); ++ if (isa(elementType)) { ++ auto nan_lhs_condition = buildNaNLhsSelectionCondition( ++ builder, loc, lhs_value, rhs_value, lt_index_pred); ++ final_lhs_condition = ++ OrOp::create(builder, loc, final_lhs_condition, nan_lhs_condition) ++ .getResult(); ++ } + +- // Select Final Results: +- // if final_lhs_condition: +- // return (lhs_value, lhs_index) +- // else: +- // return (rhs_value, rhs_index) ++ // Use MaxOp for the value so that NaNs propagate properly and unused-index ++ // reductions simplify to kMaximum. + auto selected_value = +- SelectOp::create(builder, loc, final_lhs_condition, lhs_value, rhs_value) +- .getResult(); ++ MaxOp::create(builder, loc, lhs_value, rhs_value).getResult(); + auto selected_index = + SelectOp::create(builder, loc, final_lhs_condition, lhs_index, rhs_index) + .getResult(); diff --ruN a/stablehlo/stablehlo/dialect/StablehloOps.td b/stablehlo/stablehlo/dialect/StablehloOps.td --- stablehlo/stablehlo/dialect/StablehloOps.td +++ stablehlo/stablehlo/dialect/StablehloOps.td -@@ -2525,8 +2525,9 @@ +@@ -38,6 +38,7 @@ + + let useDefaultAttributePrinterParser = 0; + let useDefaultTypePrinterParser = 0; ++ let useStrictPropertiesInAssemblyFormat = 0; + } + + class StableHLO_Op traits = []> : +@@ -2525,8 +2526,9 @@ "::mlir::TypeRange":$resultTypes, "::mlir::ValueRange":$inputs, "::mlir::ArrayRef<::mlir::NamedAttribute>":$attributes ), [{ @@ -1021,6 +1093,17 @@ diff --ruN a/stablehlo/stablehlo/dialect/VhloDialect.td b/stablehlo/stablehlo/di }]; let useDefaultAttributePrinterParser = 0; +diff --ruN a/stablehlo/stablehlo/dialect/VhloDialect.td b/stablehlo/stablehlo/dialect/VhloDialect.td +--- stablehlo/stablehlo/dialect/VhloDialect.td ++++ stablehlo/stablehlo/dialect/VhloDialect.td +@@ -63,6 +63,7 @@ + + let useDefaultAttributePrinterParser = 0; + let useDefaultTypePrinterParser = 0; ++ let useStrictPropertiesInAssemblyFormat = 0; + } + + #endif // STABLEHLO_DIALECT_VHLO_DIALECT diff --ruN a/stablehlo/stablehlo/dialect/VhloEnums.td b/stablehlo/stablehlo/dialect/VhloEnums.td --- stablehlo/stablehlo/dialect/VhloEnums.td +++ stablehlo/stablehlo/dialect/VhloEnums.td @@ -1244,6 +1327,119 @@ diff --ruN a/stablehlo/stablehlo/integrations/c/StablehloUnifiedApi.cpp b/stable resultsVec.push_back(wrap(result)); } return mlirArrayAttrGet(mlirModuleGetContext(module), resultsVec.size(), +diff --ruN a/stablehlo/stablehlo/integrations/cpp/builder/BUILD.bazel b/stablehlo/stablehlo/integrations/cpp/builder/BUILD.bazel +--- stablehlo/stablehlo/integrations/cpp/builder/BUILD.bazel ++++ stablehlo/stablehlo/integrations/cpp/builder/BUILD.bazel +@@ -228,6 +228,9 @@ + ":func_builder", + ":mlir_builder", + ":stablehlo_builder", ++ "//:reference_ops", ++ "//:reference_tensor", ++ "//:reference_value", + "//:register", + "//:stablehlo_ops", + "@llvm-project//mlir:IR", +diff --ruN a/stablehlo/stablehlo/integrations/cpp/builder/StablehloBuilderTest.cpp b/stablehlo/stablehlo/integrations/cpp/builder/StablehloBuilderTest.cpp +--- stablehlo/stablehlo/integrations/cpp/builder/StablehloBuilderTest.cpp ++++ stablehlo/stablehlo/integrations/cpp/builder/StablehloBuilderTest.cpp +@@ -15,6 +15,7 @@ + + #include + #include ++#include + #include + + #include "gtest/gtest.h" +@@ -29,6 +30,12 @@ + #include "mlir/Support/LLVM.h" + #include "stablehlo/dialect/Register.h" + #include "stablehlo/dialect/StablehloOps.h" ++#include "stablehlo/reference/Ops.h" ++#include "stablehlo/reference/Tensor.h" ++#include "stablehlo/reference/Value.h" ++ ++// Include builder headers after reference/Ops.h so StablehloBuilder.h's ++// Transpose() function does not hide the Transpose enum under -fno-modules. + #include "stablehlo/integrations/cpp/builder/AttrTypeBuilderUtil.h" + #include "stablehlo/integrations/cpp/builder/FuncBuilder.h" + #include "stablehlo/integrations/cpp/builder/MlirBuilder.h" +@@ -310,6 +317,75 @@ + EXPECT_EQ(expected, debugString(*module)); + } + ++TEST(MlirBuilderTest, BuildMaxAndArgmaxBodyUsesMaxOp) { ++ MLIRContext context; ++ context.loadDialect(); ++ OpBuilder builder(&context); ++ OwningOpRef module = ModuleOp::create(builder.getUnknownLoc()); ++ builder.setInsertionPointToEnd(module->getBody()); ++ ++ buildMaxAndArgmaxBody(builder.getF32Type(), builder.getI32Type(), ++ module->getBodyRegion(), builder); ++ ++ Operation& returnOp = module->getBody()->back(); ++ EXPECT_TRUE(isa(returnOp.getOperand(0).getDefiningOp())); ++ EXPECT_TRUE(isa(returnOp.getOperand(1).getDefiningOp())); ++} ++ ++TEST(MlirBuilderTest, BuildMaxAndArgmaxBodyPropagatesNaN) { ++ MLIRContext context; ++ context.loadDialect(); ++ OpBuilder builder(&context); ++ OwningOpRef module = ModuleOp::create(builder.getUnknownLoc()); ++ builder.setInsertionPointToEnd(module->getBody()); ++ ++ buildMaxAndArgmaxBody(builder.getF32Type(), builder.getI32Type(), ++ module->getBodyRegion(), builder); ++ ++ auto makeScalar = [&](TypedAttr attr) { ++ return InterpreterValue(makeTensor(DenseElementsAttr::get( ++ RankedTensorType::get({}, attr.getType()), attr))); ++ }; ++ float nan = std::numeric_limits::quiet_NaN(); ++ // SelectOp would drop NaN when lhs=NaN and rhs=5.0; MaxOp propagates NaN. ++ auto res = ++ eval(module->getBodyRegion(), {makeScalar(builder.getF32FloatAttr(nan)), ++ makeScalar(builder.getI32IntegerAttr(0)), ++ makeScalar(builder.getF32FloatAttr(5.0f)), ++ makeScalar(builder.getI32IntegerAttr(1))}); ++ EXPECT_TRUE(res[0].getTensor().get({}).getFloatValue().isNaN()); ++} ++ ++TEST(MlirBuilderTest, BuildMaxAndArgmaxBodySelectsNaNIndex) { ++ MLIRContext context; ++ context.loadDialect(); ++ OpBuilder builder(&context); ++ OwningOpRef module = ModuleOp::create(builder.getUnknownLoc()); ++ builder.setInsertionPointToEnd(module->getBody()); ++ ++ buildMaxAndArgmaxBody(builder.getF32Type(), builder.getI32Type(), ++ module->getBodyRegion(), builder); ++ ++ auto evalIndex = [&](float lhs_val, int32_t lhs_idx, float rhs_val, ++ int32_t rhs_idx) { ++ auto makeScalar = [&](TypedAttr attr) { ++ return InterpreterValue(makeTensor(DenseElementsAttr::get( ++ RankedTensorType::get({}, attr.getType()), attr))); ++ }; ++ auto res = eval(module->getBodyRegion(), ++ {makeScalar(builder.getF32FloatAttr(lhs_val)), ++ makeScalar(builder.getI32IntegerAttr(lhs_idx)), ++ makeScalar(builder.getF32FloatAttr(rhs_val)), ++ makeScalar(builder.getI32IntegerAttr(rhs_idx))}); ++ EXPECT_TRUE(res[0].getTensor().get({}).getFloatValue().isNaN()); ++ return res[1].getTensor().get({}).getIntegerValue().getSExtValue(); ++ }; ++ float nan = std::numeric_limits::quiet_NaN(); ++ EXPECT_EQ(evalIndex(nan, 0, 5.0f, 1), 0); ++ EXPECT_EQ(evalIndex(5.0f, 0, nan, 1), 1); ++ EXPECT_EQ(evalIndex(nan, 2, nan, 1), 1); ++} ++ + TEST(MlirBuilderTest, GatherOp) { + std::string expected = R"mlir(module { + func.func @main(%arg0: tensor<3xi64>, %arg1: tensor<1x1xi64>) -> tensor<1xi64> { diff --ruN a/stablehlo/stablehlo/integrations/python/StablehloApi.cpp b/stablehlo/stablehlo/integrations/python/StablehloApi.cpp --- stablehlo/stablehlo/integrations/python/StablehloApi.cpp +++ stablehlo/stablehlo/integrations/python/StablehloApi.cpp From a599fd9420cfc09ce11f5550108f1b90ebe78d65 Mon Sep 17 00:00:00 2001 From: Dmitri Latushko Date: Tue, 29 Sep 2026 18:41:12 -0700 Subject: [PATCH 12/21] Fix data race in ThreadSafeStatus::status() by returning absl::Status by value so it is copied while the lock is held. PiperOrigin-RevId: 990657976 --- .../batching_util/batch_resource_base.cc | 5 ++-- .../batching_util/threadsafe_status.cc | 4 +-- .../kernels/batching_util/threadsafe_status.h | 2 +- .../batching_util/threadsafe_status_test.cc | 28 +++++++++++++++++++ 4 files changed, 34 insertions(+), 5 deletions(-) diff --git a/tensorflow/core/kernels/batching_util/batch_resource_base.cc b/tensorflow/core/kernels/batching_util/batch_resource_base.cc index 33e1ce8977aacd..14e86077c34b80 100644 --- a/tensorflow/core/kernels/batching_util/batch_resource_base.cc +++ b/tensorflow/core/kernels/batching_util/batch_resource_base.cc @@ -1081,8 +1081,9 @@ absl::Status BatchResourceBase::ConcatInputTensors( // the priority queue). If so, skip the output concatenation — the // output TensorMatrix may contain uninitialized entries that would // cause a crash in Concat/memcpy. - if (!input_task->status->status().ok()) { - input_task->FinishTask(input_task->status->status()); + absl::Status task_status = input_task->status->status(); + if (!task_status.ok()) { + input_task->FinishTask(task_status); return; } OpKernelContext* context = input_task->context; diff --git a/tensorflow/core/kernels/batching_util/threadsafe_status.cc b/tensorflow/core/kernels/batching_util/threadsafe_status.cc index fc4bd4c6c8e37e..57502be2295b9a 100644 --- a/tensorflow/core/kernels/batching_util/threadsafe_status.cc +++ b/tensorflow/core/kernels/batching_util/threadsafe_status.cc @@ -21,13 +21,13 @@ limitations under the License. #include "tensorflow/core/platform/mutex.h" namespace tensorflow { -const absl::Status& ThreadSafeStatus::status() const& { +absl::Status ThreadSafeStatus::status() const& { tf_shared_lock lock(mutex_); return status_; } absl::Status ThreadSafeStatus::status() && { - tf_shared_lock lock(mutex_); + mutex_lock lock(mutex_); return std::move(status_); } diff --git a/tensorflow/core/kernels/batching_util/threadsafe_status.h b/tensorflow/core/kernels/batching_util/threadsafe_status.h index 68e94f705f0d47..9c368cff48ea4e 100644 --- a/tensorflow/core/kernels/batching_util/threadsafe_status.h +++ b/tensorflow/core/kernels/batching_util/threadsafe_status.h @@ -40,7 +40,7 @@ namespace tensorflow { // When updated in a multi-threading setup, only the first error is retained. class ThreadSafeStatus { public: - const absl::Status& status() const& TF_LOCKS_EXCLUDED(mutex_); + absl::Status status() const& TF_LOCKS_EXCLUDED(mutex_); absl::Status status() && TF_LOCKS_EXCLUDED(mutex_); // Retains the first error status: replaces the current status with diff --git a/tensorflow/core/kernels/batching_util/threadsafe_status_test.cc b/tensorflow/core/kernels/batching_util/threadsafe_status_test.cc index 834e7eb18455d0..b9acc5ef167a89 100644 --- a/tensorflow/core/kernels/batching_util/threadsafe_status_test.cc +++ b/tensorflow/core/kernels/batching_util/threadsafe_status_test.cc @@ -15,6 +15,10 @@ limitations under the License. #include "tensorflow/core/kernels/batching_util/threadsafe_status.h" +#include +#include // NOLINT(build/c++11) +#include + #include "tensorflow/core/lib/core/status_test_util.h" #include "tensorflow/core/platform/errors.h" #include "tensorflow/core/platform/test.h" @@ -47,5 +51,29 @@ TEST(ThreadSafeStatus, Move) { TF_EXPECT_OK(std::move(status).status()); } +TEST(ThreadSafeStatus, ConcurrentReadAndUpdate) { + ThreadSafeStatus status; + std::atomic done{false}; + std::thread reader([&]() { + while (!done.load(std::memory_order_relaxed)) { + absl::Status s = status.status(); + if (!s.ok()) { + EXPECT_EQ(s.code(), error::INTERNAL); + } + } + }); + + std::thread updater([&]() { + for (int i = 0; i < 1000; ++i) { + status.Update(absl::InternalError("concurrent error")); + } + done.store(true, std::memory_order_relaxed); + }); + + updater.join(); + reader.join(); + EXPECT_EQ(status.status().code(), error::INTERNAL); +} + } // namespace } // namespace tensorflow From c53bc6e4ba4c8877b6d3a5bc8842971b7196fe30 Mon Sep 17 00:00:00 2001 From: "William S. Moses" Date: Tue, 29 Sep 2026 19:15:25 -0700 Subject: [PATCH 13/21] [XLA][SPMD] Optimize 2-concatenate for multi devices PiperOrigin-RevId: 990670493 --- .../xla/xla/service/spmd/spmd_partitioner.cc | 110 +++++++++++++++- .../xla/xla/service/spmd/spmd_partitioner.h | 10 ++ .../xla/service/spmd/spmd_partitioner_test.cc | 124 ++++++++++++++++++ 3 files changed, 240 insertions(+), 4 deletions(-) diff --git a/third_party/xla/xla/service/spmd/spmd_partitioner.cc b/third_party/xla/xla/service/spmd/spmd_partitioner.cc index 30b8dfe86497e6..fdfb08bf505eeb 100644 --- a/third_party/xla/xla/service/spmd/spmd_partitioner.cc +++ b/third_party/xla/xla/service/spmd/spmd_partitioner.cc @@ -2986,6 +2986,13 @@ absl::Status SpmdPartitioningVisitor::HandleTriangularSolve( } absl::Status SpmdPartitioningVisitor::HandleConcatenate(HloInstruction* hlo) { + if (module_->config().debug_options().xla_enable_enzyme_comms_opt()) { + ABSL_ASSIGN_OR_RETURN(bool handled, + TryHandleConcatenateWithConstantOffsets(hlo)); + if (handled) { + return absl::OkStatus(); + } + } return HandleElementwiseWithDimsToReplicate(hlo, {hlo->concatenate_dimension()}); } @@ -5055,9 +5062,15 @@ SpmdPartitioningVisitor::ProcessUpdatePieceExtractOperand( bool enableBroadcastOptimization = actual_update->operand(0)->shape().dimensions().empty(); if (enableBroadcastOptimization) { + // The concatenate handler has no input tensor; the target is the + // result's shard shape either way. + Shape broadcast_shape = + input_tensor != nullptr + ? GetPartitionedHlo(input_tensor).hlo()->shape() + : MakePartitionedShape(hlo->shape(), hlo->sharding()); newOperand = add_hlo(HloInstruction::CreateBroadcast( - GetPartitionedHlo(input_tensor).hlo()->shape(), - GetPartitionedHlo(actual_update->operand(0)).hlo(), {})); + broadcast_shape, GetPartitionedHlo(actual_update->operand(0)).hlo(), + {})); newOperand->set_sharding(hlo->sharding()); } } else { @@ -5365,8 +5378,14 @@ SpmdPartitioningVisitor::ProcessUpdatePieceExtractOperand( break; } if (ShardCountAtDim(hlo->sharding(), i) > 1) { - int64_t dus_start = - dus->operand(i + 2)->literal().GetIntegralAsS64({}).value(); + // A concatenate handled through this path has no + // dynamic-update-slice; the piece's own start is what the + // index would say. + int64_t dus_start = dus != nullptr ? dus->operand(i + 2) + ->literal() + .GetIntegralAsS64({}) + .value() + : piece_dus_starts[i]; int64_t slice_start = slice->slice_starts(i); if (absl::c_linear_search(reverse_dims, i)) { slice_start = slice->operand(0)->shape().dimensions(i) - @@ -5521,6 +5540,89 @@ SpmdPartitioningVisitor::ProcessUpdatePieceExtractOperand( return newOperand; } +absl::StatusOr +SpmdPartitioningVisitor::TryHandleConcatenateWithConstantOffsets( + HloInstruction* hlo) { + const HloSharding& sharding = hlo->sharding(); + if (hlo->shape().IsTuple() || !sharding.IsTiled() || + hlo->operand_count() < 2) { + return false; + } + const int64_t dim = hlo->concatenate_dimension(); + const int64_t num_shards = sharding.dimension(dim); + if (num_shards <= 1) { + return false; + } + // Every operand has to be laid out like the result, so that the only data + // movement is the shift along the concatenate dimension. + for (const HloInstruction* operand : hlo->operands()) { + if (!operand->has_sharding() || operand->sharding() != sharding) { + return false; + } + } + + const int64_t rank = hlo->shape().dimensions().size(); + const int64_t full_size = hlo->shape().dimensions(dim); + const int64_t shard_size = CeilOfRatio(full_size, num_shards); + + // The largest operand stays where it is; the others are written in around + // it. + int64_t largest = 0; + std::vector offsets(hlo->operand_count()); + int64_t offset = 0; + for (int64_t i = 0; i < hlo->operand_count(); ++i) { + offsets[i] = offset; + offset += hlo->operand(i)->shape().dimensions(dim); + if (hlo->operand(i)->shape().dimensions(dim) > + hlo->operand(largest)->shape().dimensions(dim)) { + largest = i; + } + } + // Halo exchange only reaches the direct neighbour, so every shift has to be + // smaller than a shard: the largest operand's offset, and the size of each + // operand written in over it. + if (offsets[largest] >= shard_size) { + return false; + } + for (int64_t i = 0; i < hlo->operand_count(); ++i) { + if (i != largest && + hlo->operand(i)->shape().dimensions(dim) >= shard_size) { + return false; + } + } + + PaddingConfig padding_config; + for (int64_t d = 0; d < rank; ++d) { + auto* padding_dim = padding_config.add_dimensions(); + padding_dim->set_interior_padding(0); + padding_dim->set_edge_padding_low(d == dim ? offsets[largest] : 0); + padding_dim->set_edge_padding_high( + d == dim ? full_size - offsets[largest] - + hlo->operand(largest)->shape().dimensions(dim) + : 0); + } + HloInstruction* zero = b_.AddInstruction(HloInstruction::CreateConstant( + LiteralUtil::Zero(hlo->shape().element_type()))); + HloInstruction* current = + PadHelper(*this, GetPartitionedHlo(hlo->operand(largest)), zero, + padding_config, hlo->shape(), sharding); + if (current == nullptr) { + return false; + } + for (int64_t i = 0; i < hlo->operand_count(); ++i) { + if (i == largest) { + continue; + } + std::vector starts(rank, 0); + starts[dim] = offsets[i]; + ABSL_ASSIGN_OR_RETURN(current, + ProcessUpdatePiece(hlo, /*input_tensor=*/nullptr, + hlo->operand(i), starts, current)); + } + SetPartitionedHlo(hlo, current); + return true; +} + absl::Status SpmdPartitioningVisitor::HandleDUSAllPartitionedSliceDimsHaveConstantIndices( HloInstruction* hlo, const HloInstruction* input_tensor, diff --git a/third_party/xla/xla/service/spmd/spmd_partitioner.h b/third_party/xla/xla/service/spmd/spmd_partitioner.h index 3aa388b8e1dc97..7d4e4376c060c2 100644 --- a/third_party/xla/xla/service/spmd/spmd_partitioner.h +++ b/third_party/xla/xla/service/spmd/spmd_partitioner.h @@ -963,6 +963,16 @@ class SpmdPartitioningVisitor : public DfsHloVisitorWithDefault { std::vector partitioned_slice_dims); // Method 3: All partitioned slice dimensions have compile-time constant // indices. + // Partitions a concatenate along a partitioned dimension without + // replicating that dimension: the largest operand is padded into its place + // in the result, and each remaining operand is written in at its constant + // offset the way a dynamic-update-slice with constant indices is, so data + // only moves between neighbouring shards. Returns false, having done + // nothing, when the concatenate does not fit the shapes this handles. Only + // used with xla_enable_enzyme_comms_opt. + absl::StatusOr TryHandleConcatenateWithConstantOffsets( + HloInstruction* hlo); + absl::Status HandleDUSAllPartitionedSliceDimsHaveConstantIndices( HloInstruction* hlo, const HloInstruction* input_tensor, const HloInstruction* update_tensor); diff --git a/third_party/xla/xla/service/spmd/spmd_partitioner_test.cc b/third_party/xla/xla/service/spmd/spmd_partitioner_test.cc index bbb367d23f79a2..7b46d91f3b3cf7 100644 --- a/third_party/xla/xla/service/spmd/spmd_partitioner_test.cc +++ b/third_party/xla/xla/service/spmd/spmd_partitioner_test.cc @@ -19630,6 +19630,130 @@ ENTRY entry { EXPECT_TRUE(has_scatter); } +// A row spliced in front of a tensor along a dimension partitioned two ways. +// Replicating the concatenate dimension to do this is an all-to-all on a 2x2 +// mesh; with the enzyme comms opt the row is written in at its offset instead +// and only the shard boundary moves, between neighbours. +TEST_P(SpmdPartitioningTest, ConcatenateAlongPartitionedDimWithEnzymeOpt) { + absl::string_view hlo_string = R"( +HloModule module + +ENTRY entry { + %row = f32[4,1,8] parameter(0), sharding={devices=[1,2,2]<=[4]} + %x = f32[4,7,8] parameter(1), sharding={devices=[1,2,2]<=[4]} + ROOT %concat = f32[4,8,8] concatenate(%row, %x), dimensions={1}, + sharding={devices=[1,2,2]<=[4]} +})"; + ASSERT_OK_AND_ASSIGN(auto module, + PartitionComputation(hlo_string, /*num_devices=*/4, + SpmdPartitionerOptions(), + /*enable_enzyme_opt=*/true)); + VLOG(1) << module->ToString(); + const HloComputation* entry = module->entry_computation(); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllToAll), 0); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllGather), 0); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllReduce), 0); + EXPECT_GT(NumOfInstructions(entry, HloOpcode::kCollectivePermute), 0); + EXPECT_THAT(entry->root_instruction(), op::Shape("f32[4,4,4]")); +} + +// The same with the row at the end, the other form the algebraic simplifier +// produces from a dynamic-update-slice into a pad. +TEST_P(SpmdPartitioningTest, + ConcatenateAlongPartitionedDimTrailingOperandWithEnzymeOpt) { + absl::string_view hlo_string = R"( +HloModule module + +ENTRY entry { + %x = f32[4,7,8] parameter(0), sharding={devices=[1,2,2]<=[4]} + %row = f32[4,1,8] parameter(1), sharding={devices=[1,2,2]<=[4]} + ROOT %concat = f32[4,8,8] concatenate(%x, %row), dimensions={1}, + sharding={devices=[1,2,2]<=[4]} +})"; + ASSERT_OK_AND_ASSIGN(auto module, + PartitionComputation(hlo_string, /*num_devices=*/4, + SpmdPartitionerOptions(), + /*enable_enzyme_opt=*/true)); + VLOG(1) << module->ToString(); + const HloComputation* entry = module->entry_computation(); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllToAll), 0); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllGather), 0); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllReduce), 0); + EXPECT_GT(NumOfInstructions(entry, HloOpcode::kCollectivePermute), 0); + EXPECT_THAT(entry->root_instruction(), op::Shape("f32[4,4,4]")); +} + +// Operands laid out differently from the result are left to the default +// handling. +TEST_P(SpmdPartitioningTest, + ConcatenateAlongPartitionedDimMismatchedOperandShardingWithEnzymeOpt) { + absl::string_view hlo_string = R"( +HloModule module + +ENTRY entry { + %row = f32[4,1,8] parameter(0), sharding={devices=[1,2,2]<=[4]} + %x = f32[4,7,8] parameter(1), sharding={devices=[2,1,2]<=[4]} + ROOT %concat = f32[4,8,8] concatenate(%row, %x), dimensions={1}, + sharding={devices=[1,2,2]<=[4]} +})"; + ASSERT_OK_AND_ASSIGN(auto module, + PartitionComputation(hlo_string, /*num_devices=*/4, + SpmdPartitionerOptions(), + /*enable_enzyme_opt=*/true)); + EXPECT_THAT(module->entry_computation()->root_instruction(), + op::Shape("f32[4,4,4]")); +} + +// The device order GB-25 actually uses: the tile assignment is transposed. +TEST_P(SpmdPartitioningTest, + ConcatenateAlongPartitionedDimTransposedDevicesWithEnzymeOpt) { + absl::string_view hlo_string = R"( +HloModule module + +ENTRY entry { + %row = f64[4,1,8] parameter(0), sharding={devices=[1,2,2]<=[2,2]T(1,0)} + %x = f64[4,7,8] parameter(1), sharding={devices=[1,2,2]<=[2,2]T(1,0)} + ROOT %concat = f64[4,8,8] concatenate(%row, %x), dimensions={1}, + sharding={devices=[1,2,2]<=[2,2]T(1,0)} +})"; + ASSERT_OK_AND_ASSIGN(auto module, + PartitionComputation(hlo_string, /*num_devices=*/4, + SpmdPartitionerOptions(), + /*enable_enzyme_opt=*/true)); + VLOG(1) << module->ToString(); + const HloComputation* entry = module->entry_computation(); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllToAll), 0); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllGather), 0); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllReduce), 0); + EXPECT_GT(NumOfInstructions(entry, HloOpcode::kCollectivePermute), 0); + EXPECT_THAT(entry->root_instruction(), op::Shape("f64[4,4,4]")); +} + +// More than one operand written in around the largest. +TEST_P(SpmdPartitioningTest, + ConcatenateAlongPartitionedDimThreeOperandsWithEnzymeOpt) { + absl::string_view hlo_string = R"( +HloModule module + +ENTRY entry { + %lo = f32[4,1,8] parameter(0), sharding={devices=[1,2,2]<=[4]} + %x = f32[4,6,8] parameter(1), sharding={devices=[1,2,2]<=[4]} + %hi = f32[4,1,8] parameter(2), sharding={devices=[1,2,2]<=[4]} + ROOT %concat = f32[4,8,8] concatenate(%lo, %x, %hi), dimensions={1}, + sharding={devices=[1,2,2]<=[4]} +})"; + ASSERT_OK_AND_ASSIGN(auto module, + PartitionComputation(hlo_string, /*num_devices=*/4, + SpmdPartitionerOptions(), + /*enable_enzyme_opt=*/true)); + VLOG(1) << module->ToString(); + const HloComputation* entry = module->entry_computation(); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllToAll), 0); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllGather), 0); + EXPECT_EQ(NumOfInstructions(entry, HloOpcode::kAllReduce), 0); + EXPECT_THAT(entry->root_instruction(), op::Shape("f32[4,4,4]")); +} + } // namespace } // namespace spmd } // namespace xla From 51e5cf1197b166e4ad33abab194f3e1652de2ee1 Mon Sep 17 00:00:00 2001 From: Matthias Guenther Date: Tue, 29 Sep 2026 19:27:58 -0700 Subject: [PATCH 14/21] Add llvm_xz dependency and BUILD definition. Adds the llvm_xz (v5.8.3) external archive in WORKSPACE.bazel and provides third_party/xz.BUILD to build the liblzma library. PiperOrigin-RevId: 990674437 --- .../xla/third_party/stablehlo/temporary.patch | 30 ++++++++++++++----- 1 file changed, 23 insertions(+), 7 deletions(-) diff --git a/third_party/xla/third_party/stablehlo/temporary.patch b/third_party/xla/third_party/stablehlo/temporary.patch index 6a05d2e1510d36..678053cd922786 100644 --- a/third_party/xla/third_party/stablehlo/temporary.patch +++ b/third_party/xla/third_party/stablehlo/temporary.patch @@ -36,6 +36,28 @@ diff --ruN a/stablehlo/BUILD.bazel b/stablehlo/BUILD.bazel +# deprecation = "Use //stablehlo/tests:test_utils instead.", +# ) +# copybara:uncomment_end +diff --ruN a/stablehlo/WORKSPACE.bazel b/stablehlo/WORKSPACE.bazel +--- stablehlo/WORKSPACE.bazel ++++ stablehlo/WORKSPACE.bazel +@@ -119,6 +119,18 @@ + ], + ) + ++# Required by the LLVM Bazel overlay (`@llvm-project//third-party:lzma`). ++http_archive( ++ name = "llvm_xz", ++ build_file = "//third_party:xz.BUILD", ++ sha256 = "3d3a1b973af218114f4f889bbaa2f4c037deaae0c8e815eec381c3d546b974a0", ++ strip_prefix = "xz-5.8.3", ++ urls = [ ++ "https://storage.googleapis.com/mirror.tensorflow.org/github.com/tukaani-project/xz/releases/download/v5.8.3/xz-5.8.3.tar.gz", ++ "https://github.com/tukaani-project/xz/releases/download/v5.8.3/xz-5.8.3.tar.gz", ++ ], ++) ++ + # These need to come in this specific order or else Bazel will complain about + # missing/circular dependencies. + load("//third_party/llvm:workspace.bzl", llvm = "repo") diff --ruN a/stablehlo/build_tools/lit_test_suite.bzl b/stablehlo/build_tools/lit_test_suite.bzl --- stablehlo/build_tools/lit_test_suite.bzl +++ stablehlo/build_tools/lit_test_suite.bzl @@ -1085,19 +1107,13 @@ diff --ruN a/stablehlo/stablehlo/dialect/VhloBytecode.h b/stablehlo/stablehlo/di diff --ruN a/stablehlo/stablehlo/dialect/VhloDialect.td b/stablehlo/stablehlo/dialect/VhloDialect.td --- stablehlo/stablehlo/dialect/VhloDialect.td +++ stablehlo/stablehlo/dialect/VhloDialect.td -@@ -59,6 +59,7 @@ +@@ -59,10 +59,12 @@ 1.18.0: Add `result_tilings` attribute to `custom_call` op. 1.19.0: Add CollectiveReduceOp. 1.20.0: Add has_dynamic_root to collective_broadcast op. Allow mixed fp8 operands in `convolution` and `dynamic_conv` ops. + 1.21.0: Add F8E4M3FN_F8E4M3FN_F32_X3 and F8E4M3FN_F8E4M3FN_F32_X4 dot algorithms. }]; - let useDefaultAttributePrinterParser = 0; -diff --ruN a/stablehlo/stablehlo/dialect/VhloDialect.td b/stablehlo/stablehlo/dialect/VhloDialect.td ---- stablehlo/stablehlo/dialect/VhloDialect.td -+++ stablehlo/stablehlo/dialect/VhloDialect.td -@@ -63,6 +63,7 @@ - let useDefaultAttributePrinterParser = 0; let useDefaultTypePrinterParser = 0; + let useStrictPropertiesInAssemblyFormat = 0; From 2091fe07b0092094f2bc32e6e516ca1ebeece739 Mon Sep 17 00:00:00 2001 From: Volodymyr Kysenko Date: Tue, 29 Sep 2026 20:04:02 -0700 Subject: [PATCH 15/21] Support broadcast-based grouped-query attention (GQA) during prefill in YNNPACK. Instead of folding query heads into the row/sequence dimension for GQA during prefill, split Q into [n_kv, g] and expand K, V, and the mask to 5D so they broadcast across the group dimension. This updates the K and V transposes to 5D, avoids splitting and fusing logits around mask addition, and fuses the [n_kv, g] head dimensions back together after the P @ V matmul. GQA folding is retained for decode paths and sequence-major inputs. PiperOrigin-RevId: 990686744 --- .../lite/delegates/ynnpack/attention.cc | 88 +++++++++++++------ .../lite/delegates/ynnpack/attention_test.cc | 4 +- 2 files changed, 65 insertions(+), 27 deletions(-) diff --git a/tensorflow/lite/delegates/ynnpack/attention.cc b/tensorflow/lite/delegates/ynnpack/attention.cc index eb54c9c6de4f46..f7c50524f95a69 100644 --- a/tensorflow/lite/delegates/ynnpack/attention.cc +++ b/tensorflow/lite/delegates/ynnpack/attention.cc @@ -350,7 +350,16 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, TF_LITE_ENSURE(context, n_kv > 0 && n_q % n_kv == 0); const size_t g_heads_per_kv = static_cast(n_q / n_kv); - if (g_heads_per_kv > 1) { + const int q_seq_dim = is_seq_major ? 1 : 2; + bool use_decode1 = + !is_seq_major && + (q_tensor.dims->data[q_seq_dim] * static_cast(g_heads_per_kv) <= 32); + // Prefill with grouped heads: keep the head axis as [n_kv, g] and let the dot + // broadcast K/V over g, instead of folding g into the row axis. + const bool gqa_batch = g_heads_per_kv > 1 && !is_seq_major && !use_decode1; + const bool gqa_fold = g_heads_per_kv > 1 && !gqa_batch; + + if (gqa_fold) { uint32_t q_5d_id = YNN_INVALID_VALUE_ID; const size_t q_splits[2] = {static_cast(n_kv), g_heads_per_kv}; TF_LITE_ENSURE_YNN_STATUS(ynn_define_split_dim(subgraph, /*axis=*/1, @@ -362,10 +371,19 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, q_trans_id = q_packed_id; } - const int q_seq_dim = is_seq_major ? 1 : 2; - bool use_decode1 = - !is_seq_major && - (q_tensor.dims->data[q_seq_dim] * static_cast(g_heads_per_kv) <= 32); + if (gqa_batch) { + const size_t q_splits[2] = {static_cast(n_kv), g_heads_per_kv}; + uint32_t q_5d_id = YNN_INVALID_VALUE_ID; + TF_LITE_ENSURE_YNN_STATUS(ynn_define_split_dim(subgraph, /*axis=*/1, + /*num_splits=*/2, q_splits, + q_trans_id, &q_5d_id, 0)); + q_trans_id = q_5d_id; + const int32_t expand_axis = 2; + uint32_t k_5d_id = YNN_INVALID_VALUE_ID; + TF_LITE_ENSURE_YNN_STATUS(ynn_define_static_expand_dims( + subgraph, /*num_new_axes=*/1, &expand_axis, k_trans_id, &k_5d_id, 0)); + k_trans_id = k_5d_id; + } bool need_slice_out = false; uint32_t post_bmm_id = YNN_INVALID_VALUE_ID; @@ -382,13 +400,14 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, &q_scaled_id, 0)); // Scores: S = Q @ K^T, [B, H, Q, S]. - const int32_t swap_last_two_perm[] = {0, 1, 3, 2}; + const int32_t swap_last_two_perm[] = {-1, -2}; uint32_t scores_id = YNN_INVALID_VALUE_ID; if (use_decode1) { // Compute S^T = K @ Q^T and transpose the (small) result. uint32_t q_scaled_t_id = YNN_INVALID_VALUE_ID; TF_LITE_ENSURE_YNN_STATUS(ynn_define_static_transpose( - subgraph, 4, swap_last_two_perm, q_scaled_id, &q_scaled_t_id, 0)); + subgraph, 2, swap_last_two_perm, q_scaled_id, &q_scaled_t_id, + YNN_NODE_FLAG_KEEP_DIMS)); uint32_t scores_ts_id = YNN_INVALID_VALUE_ID; TF_LITE_ENSURE_YNN_STATUS( @@ -396,11 +415,13 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, YNN_INVALID_VALUE_ID, &scores_ts_id, 0)); TF_LITE_ENSURE_YNN_STATUS(ynn_define_static_transpose( - subgraph, 4, swap_last_two_perm, scores_ts_id, &scores_id, 0)); + subgraph, 2, swap_last_two_perm, scores_ts_id, &scores_id, + YNN_NODE_FLAG_KEEP_DIMS)); } else { uint32_t k_trans_t_id = YNN_INVALID_VALUE_ID; - TF_LITE_ENSURE_YNN_STATUS(ynn_define_static_transpose( - subgraph, 4, swap_last_two_perm, k_trans_id, &k_trans_t_id, 0)); + TF_LITE_ENSURE_YNN_STATUS( + ynn_define_static_transpose(subgraph, 2, swap_last_two_perm, k_trans_id, + &k_trans_t_id, YNN_NODE_FLAG_KEEP_DIMS)); TF_LITE_ENSURE_YNN_STATUS( ynn_define_dot(subgraph, /*num_k_dims=*/1, q_scaled_id, k_trans_t_id, @@ -445,22 +466,24 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, &sliced_mask_id, /*flags=*/0)); mask_to_add_id = sliced_mask_id; } - if (g_heads_per_kv > 1) { - const size_t logits_splits[2] = {g_heads_per_kv, 0}; - uint32_t logits_5d_id = YNN_INVALID_VALUE_ID; - TF_LITE_ENSURE_YNN_STATUS( - ynn_define_split_dim(subgraph, /*axis=*/2, /*num_splits=*/2, - logits_splits, logits_id, &logits_5d_id, 0)); - + if (gqa_fold || gqa_batch) { const int32_t expand_axis = 2; uint32_t mask_5d_id = YNN_INVALID_VALUE_ID; TF_LITE_ENSURE_YNN_STATUS(ynn_define_static_expand_dims( subgraph, /*num_new_axes=*/1, &expand_axis, mask_to_add_id, &mask_5d_id, 0)); + mask_to_add_id = mask_5d_id; + } + if (gqa_fold) { + const size_t logits_splits[2] = {g_heads_per_kv, 0}; + uint32_t logits_5d_id = YNN_INVALID_VALUE_ID; + TF_LITE_ENSURE_YNN_STATUS( + ynn_define_split_dim(subgraph, /*axis=*/2, /*num_splits=*/2, + logits_splits, logits_id, &logits_5d_id, 0)); uint32_t masked_logits_5d_id = YNN_INVALID_VALUE_ID; TF_LITE_ENSURE_YNN_STATUS(ynn_define_binary(subgraph, ynn_binary_add, - logits_5d_id, mask_5d_id, + logits_5d_id, mask_to_add_id, &masked_logits_5d_id, 0)); TF_LITE_ENSURE_YNN_STATUS( @@ -493,6 +516,14 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, &sliced_v_val_id, /*flags=*/0)); current_v_val_id = sliced_v_val_id; } + if (gqa_batch) { + const int32_t expand_axis = 2; + uint32_t v_5d_id = YNN_INVALID_VALUE_ID; + TF_LITE_ENSURE_YNN_STATUS(ynn_define_static_expand_dims( + subgraph, /*num_new_axes=*/1, &expand_axis, current_v_val_id, &v_5d_id, + 0)); + current_v_val_id = v_5d_id; + } // O = P @ V. if (use_decode1 && !is_seq_major) { @@ -501,8 +532,9 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, // V is [B, N, H, S] // V @ P^T is [B, N, H, 1] -> transpose to [B, N, 1, H] uint32_t probs_t_id = YNN_INVALID_VALUE_ID; - TF_LITE_ENSURE_YNN_STATUS(ynn_define_static_transpose( - subgraph, 4, swap_last_two_perm, probs_id, &probs_t_id, 0)); + TF_LITE_ENSURE_YNN_STATUS( + ynn_define_static_transpose(subgraph, 2, swap_last_two_perm, probs_id, + &probs_t_id, YNN_NODE_FLAG_KEEP_DIMS)); uint32_t post_bmm_t_id = YNN_INVALID_VALUE_ID; TF_LITE_ENSURE_YNN_STATUS( @@ -510,14 +542,15 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, YNN_INVALID_VALUE_ID, &post_bmm_t_id, 0)); TF_LITE_ENSURE_YNN_STATUS(ynn_define_static_transpose( - subgraph, 4, swap_last_two_perm, post_bmm_t_id, post_bmm_ptr, 0)); + subgraph, 2, swap_last_two_perm, post_bmm_t_id, post_bmm_ptr, + YNN_NODE_FLAG_KEEP_DIMS)); } else { // Bring V to [B, H, S, D]. - const int32_t seq_major_v_perm[] = {0, 2, 1, 3}; + const int32_t seq_major_v_perm[] = {2, 1}; uint32_t v_trans_id = YNN_INVALID_VALUE_ID; TF_LITE_ENSURE_YNN_STATUS(ynn_define_static_transpose( - subgraph, 4, is_seq_major ? seq_major_v_perm : swap_last_two_perm, - current_v_val_id, &v_trans_id, 0)); + subgraph, 2, is_seq_major ? seq_major_v_perm : swap_last_two_perm, + current_v_val_id, &v_trans_id, YNN_NODE_FLAG_KEEP_DIMS)); TF_LITE_ENSURE_YNN_STATUS( ynn_define_dot(subgraph, /*num_k_dims=*/1, probs_id, v_trans_id, @@ -527,7 +560,7 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, uint32_t post_trans_id = *post_bmm_ptr; uint32_t* post_trans_ptr = &post_trans_id; - if (g_heads_per_kv > 1) { + if (gqa_fold) { const size_t out_splits[2] = {g_heads_per_kv, 0}; uint32_t out_5d_id = YNN_INVALID_VALUE_ID; TF_LITE_ENSURE_YNN_STATUS( @@ -548,6 +581,11 @@ TfLiteStatus DefineSdpaNode(TfLiteContext* context, ynn_subgraph_t subgraph, /*axes_count=*/2, out_5d_id, post_trans_ptr, 0)); } + } else if (gqa_batch) { + post_trans_ptr = need_slice_out ? &post_trans_id : &output_val_id; + TF_LITE_ENSURE_YNN_STATUS( + ynn_define_fuse_dim(subgraph, /*axis=*/1, /*axes_count=*/2, + *post_bmm_ptr, post_trans_ptr, 0)); } else if (is_seq_major) { if (!need_slice_out) { post_trans_ptr = &output_val_id; diff --git a/tensorflow/lite/delegates/ynnpack/attention_test.cc b/tensorflow/lite/delegates/ynnpack/attention_test.cc index 033b64653b41f6..3404795aec0149 100644 --- a/tensorflow/lite/delegates/ynnpack/attention_test.cc +++ b/tensorflow/lite/delegates/ynnpack/attention_test.cc @@ -271,10 +271,10 @@ std::string PrintAttentionImplName( } TEST(AttentionGqaTest, OdmlSdpaTransposedGqaDecodeAndPrefill) { - for (int t : {1, 4, 12}) { + for (int t : {1, 4, 12, 20}) { for (int n_kv : {1, 2}) { const int b = 1; - const int s = 16; + const int s = 24; const int h = 16; const int n_q = 4; const int s_active = 11; From 276976be2f3d5eb15b593ebed94da189c047eaa3 Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 20:58:49 -0700 Subject: [PATCH 16/21] Automated Code Change PiperOrigin-RevId: 990706829 --- tensorflow/core/kernels/conv_ops_fused_impl.h | 6 +++--- tensorflow/core/kernels/conv_ops_gpu.h | 2 +- tensorflow/core/kernels/conv_ops_test.cc | 6 +++--- tensorflow/core/kernels/cudnn_rnn_ops.cc | 4 ++-- tensorflow/core/kernels/cwise_op_leakyrelu.cc | 2 +- tensorflow/core/kernels/cwise_ops.h | 14 ++++++-------- tensorflow/core/kernels/cwise_ops_common.h | 4 ++-- 7 files changed, 18 insertions(+), 20 deletions(-) diff --git a/tensorflow/core/kernels/conv_ops_fused_impl.h b/tensorflow/core/kernels/conv_ops_fused_impl.h index b7f9f08eafa7bb..53e6804f015308 100644 --- a/tensorflow/core/kernels/conv_ops_fused_impl.h +++ b/tensorflow/core/kernels/conv_ops_fused_impl.h @@ -729,7 +729,7 @@ class FusedConv2DOp : public OpKernel { using FCT = FusedComputationType; std::vector patterns; - if (std::is_same::value) { + if (std::is_same_v) { patterns = { {FCT::kBiasAdd, {"BiasAdd"}}, {FCT::kBiasAddWithRelu, {"BiasAdd", "Relu"}}, @@ -748,8 +748,8 @@ class FusedConv2DOp : public OpKernel { // identity activation function, it in theory should allow to fuse // convolution with BiasAdd, but in practice it doesn't work, cuDNN ignores // this parameter and always does Relu activation. - if (std::is_same::value) { - if (std::is_same::value || std::is_same::value) { + if (std::is_same_v) { + if (std::is_same_v || std::is_same_v) { patterns = {{FCT::kBiasAdd, {"BiasAdd"}}, {FCT::kBiasAddWithRelu, {"BiasAdd", "Relu"}}}; } else { diff --git a/tensorflow/core/kernels/conv_ops_gpu.h b/tensorflow/core/kernels/conv_ops_gpu.h index 8b02fc80f37990..c52079ad0ec305 100644 --- a/tensorflow/core/kernels/conv_ops_gpu.h +++ b/tensorflow/core/kernels/conv_ops_gpu.h @@ -64,7 +64,7 @@ int64_t GetDnnWorkspaceLimitOrDefault(); // the kernel finishes. class DnnScratchAllocator : public se::ScratchAllocator { public: - virtual ~DnnScratchAllocator() {} + ~DnnScratchAllocator() override {} DnnScratchAllocator(int64_t memory_limit, OpKernelContext* context) : memory_limit_(memory_limit), total_byte_size_(0), context_(context) {} int64_t GetMemoryLimitInBytes() override { return memory_limit_; } diff --git a/tensorflow/core/kernels/conv_ops_test.cc b/tensorflow/core/kernels/conv_ops_test.cc index 28713763cba8a3..a27ab6f772ff31 100644 --- a/tensorflow/core/kernels/conv_ops_test.cc +++ b/tensorflow/core/kernels/conv_ops_test.cc @@ -511,7 +511,7 @@ class FusedConv2DOpTest : public OpsTestBase { static constexpr int kImageBatchCount = 8; static constexpr bool kIsInt8 = - std::is_same::value || std::is_same::value; + std::is_same_v || std::is_same_v; using BiasAddGraphRunner = std::function::value || std::is_same::value; + std::is_same_v || std::is_same_v; if (exact_match) { test::ExpectEqual(x, y); } else { @@ -926,7 +926,7 @@ class FusedConv2DOpTest : public OpsTestBase { constexpr int int8_scale = 80; - using ConvT = typename std::conditional::type; + using ConvT = std::conditional_t; DataType dtype_conv = DataTypeToEnum::v(); TensorShape image_shape{image_batch_count, image_height, image_width, diff --git a/tensorflow/core/kernels/cudnn_rnn_ops.cc b/tensorflow/core/kernels/cudnn_rnn_ops.cc index bdac5a34045662..ede3cd1f8fddd0 100644 --- a/tensorflow/core/kernels/cudnn_rnn_ops.cc +++ b/tensorflow/core/kernels/cudnn_rnn_ops.cc @@ -421,7 +421,7 @@ class CudnnRnnAllocatorInTemp : public ScratchAllocator { template class CudnnRnnAllocatorInOutput : public ScratchAllocator { public: - ~CudnnRnnAllocatorInOutput() override {} + ~CudnnRnnAllocatorInOutput() override = default; CudnnRnnAllocatorInOutput(OpKernelContext* context, int output_index) : context_(context), output_index_(output_index) {} int64_t GetMemoryLimitInBytes() override { @@ -464,7 +464,7 @@ class CudnnRNNSpaceAllocator : public ScratchAllocator { explicit CudnnRNNSpaceAllocator(OpKernelContext* context) : context_(context) {} - ~CudnnRNNSpaceAllocator() override {} + ~CudnnRNNSpaceAllocator() override = default; int64_t GetMemoryLimitInBytes() override { return std::numeric_limits::max(); diff --git a/tensorflow/core/kernels/cwise_op_leakyrelu.cc b/tensorflow/core/kernels/cwise_op_leakyrelu.cc index 7bff5c9f61c1a1..37cc7187dc7352 100644 --- a/tensorflow/core/kernels/cwise_op_leakyrelu.cc +++ b/tensorflow/core/kernels/cwise_op_leakyrelu.cc @@ -92,7 +92,7 @@ class LeakyReluOp : public OpKernel { void Compute(OpKernelContext* ctx) override { const Tensor& inp = ctx->input(0); Tensor* out = nullptr; - if (std::is_same::value) { + if (std::is_same_v) { OP_REQUIRES_OK(ctx, ctx->forward_input_or_allocate_output( {0}, 0, inp.shape(), &out)); } else { diff --git a/tensorflow/core/kernels/cwise_ops.h b/tensorflow/core/kernels/cwise_ops.h index 5e7be3ef285e46..a899528bf9f346 100644 --- a/tensorflow/core/kernels/cwise_ops.h +++ b/tensorflow/core/kernels/cwise_ops.h @@ -52,9 +52,8 @@ struct scalar_arg_op> { template struct safe_scalar_binary_pow_op { - static_assert(std::is_integral::value, "Integer type expected"); - static_assert(std::is_integral::value && - std::is_signed::value, + static_assert(std::is_integral_v, "Integer type expected"); + static_assert(std::is_integral_v && std::is_signed_v, "Signed integer type expected"); bool* const error; @@ -81,7 +80,7 @@ struct functor_traits> { template struct safe_div_or_mod_op { - static_assert(std::is_integral::value, "Integer type expected"); + static_assert(std::is_integral_v, "Integer type expected"); bool* const error; @@ -377,8 +376,7 @@ struct google_floor_div { }; template -struct google_floor_div< - T, typename std::enable_if::value>::type> { +struct google_floor_div>> { EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T operator()(const T& x, const T& y) const { return x / y; @@ -1255,7 +1253,7 @@ struct safe_pow : base> { // completes; this functor must still produce a defined result on the device. template struct safe_pow_ignore_error_op { - static_assert(std::is_integral::value, "Integer type expected"); + static_assert(std::is_integral_v, "Integer type expected"); EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T operator()(const T& x, const T& y) const { if (TF_PREDICT_FALSE(y < 0)) { @@ -1372,7 +1370,7 @@ struct left_shift_op { } else if (y_clamped > sizeof(T) * CHAR_BIT - 1) { y_clamped = sizeof(T) * CHAR_BIT - 1; } - using U = typename std::make_unsigned::type; + using U = std::make_unsigned_t; return static_cast(static_cast(x) << static_cast(y_clamped)); } }; diff --git a/tensorflow/core/kernels/cwise_ops_common.h b/tensorflow/core/kernels/cwise_ops_common.h index 1c92d77e761dc6..c282bc5df02b8e 100644 --- a/tensorflow/core/kernels/cwise_ops_common.h +++ b/tensorflow/core/kernels/cwise_ops_common.h @@ -285,7 +285,7 @@ class SimpleBinaryOp : public OpKernel { const Device& eigen_device = ctx->eigen_device(); Tensor* out = nullptr; - if (std::is_same::value) { + if (std::is_same_v) { OP_REQUIRES_OK(ctx, ctx->forward_input_or_allocate_output( {0, 1}, 0, in0.shape(), &out)); } else { @@ -316,7 +316,7 @@ class UnaryOp : public OpKernel { void Compute(OpKernelContext* ctx) override { const Tensor& inp = ctx->input(0); Tensor* out = nullptr; - if (std::is_same::value) { + if (std::is_same_v) { OP_REQUIRES_OK(ctx, ctx->forward_input_or_allocate_output( {0}, 0, inp.shape(), &out)); } else { From 0fa5880f274652a1379bc63df6425955deb585e5 Mon Sep 17 00:00:00 2001 From: Dmitri Latushko Date: Tue, 29 Sep 2026 21:52:38 -0700 Subject: [PATCH 17/21] Don't rewrite integer Pow(x, -1) to Reciprocal in Grappler, so integer Pow keeps raising InvalidArgumentError for negative exponents. PiperOrigin-RevId: 990728125 --- .../optimizers/arithmetic_optimizer.cc | 4 ++- .../optimizers/arithmetic_optimizer_test.cc | 33 +++++++++++++++++++ .../math_ops/cwise_ops_binary_test.py | 20 +++++++++++ 3 files changed, 56 insertions(+), 1 deletion(-) diff --git a/tensorflow/core/grappler/optimizers/arithmetic_optimizer.cc b/tensorflow/core/grappler/optimizers/arithmetic_optimizer.cc index af2118f07fea8f..4c17495dbb65f2 100644 --- a/tensorflow/core/grappler/optimizers/arithmetic_optimizer.cc +++ b/tensorflow/core/grappler/optimizers/arithmetic_optimizer.cc @@ -3309,7 +3309,9 @@ class ConvertPowStage : public ArithmeticOptimizerStage { node->set_input(1, AsControlDependency(y->name())); AddToOptimizationQueue(node); AddToOptimizationQueue(y); - } else if (curr == complex128(-1, 0)) { + } else if (curr == complex128(-1, 0) && !DataTypeIsInteger(pow.dtype())) { + // Integer Pow rejects negative exponents, while integer Reciprocal + // computes 1 / x, so the rewrite is only valid for non-integer types. node->set_op("Reciprocal"); node->set_input(1, AsControlDependency(y->name())); AddToOptimizationQueue(node); diff --git a/tensorflow/core/grappler/optimizers/arithmetic_optimizer_test.cc b/tensorflow/core/grappler/optimizers/arithmetic_optimizer_test.cc index 7c81739d58bf15..4da6a9544b93e7 100644 --- a/tensorflow/core/grappler/optimizers/arithmetic_optimizer_test.cc +++ b/tensorflow/core/grappler/optimizers/arithmetic_optimizer_test.cc @@ -16,11 +16,13 @@ limitations under the License. #include "tensorflow/core/grappler/optimizers/arithmetic_optimizer.h" #include +#include #include "absl/strings/match.h" #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" #include "tensorflow/cc/ops/array_ops.h" +#include "tensorflow/cc/ops/const_op.h" #include "tensorflow/cc/ops/math_ops.h" #include "tensorflow/cc/ops/nn_ops.h" #include "tensorflow/cc/ops/resource_variable_ops.h" @@ -3280,6 +3282,37 @@ TEST_F(ArithmeticOptimizerTest, ConvertPow) { CompareGraphs(want, got); } +TEST_F(ArithmeticOptimizerTest, ConvertPowDoesNotRewriteIntegerNegativeOne) { + // Integer Pow rejects negative exponents, so Pow(x, -1) must not become + // Reciprocal(x), which would silently compute 1 / x instead. + tensorflow::Scope s = tensorflow::Scope::NewRootScope(); + auto x32 = ops::Const(s.WithOpName("x32"), {-1, 1}, {2}); + auto y32 = ops::Const(s.WithOpName("y32"), -1); + auto x64 = ops::Const(s.WithOpName("x64"), {-1, 1}, {2}); + auto y64 = ops::Const(s.WithOpName("y64"), {-1, -1}, {2}); + Output out32 = ops::Pow(s.WithOpName("out32"), x32, y32); + Output out64 = ops::Pow(s.WithOpName("out64"), x64, y64); + + GrapplerItem item; + item.fetch = {"out32", "out64"}; + TF_CHECK_OK(s.ToGraphDef(&item.graph)); + + GraphDef got; + ArithmeticOptimizer optimizer; + EnableOnlyConvertPow(&optimizer); + OptimizeAndPrune(&optimizer, &item, &got); + + GraphDef want; + AddNode("x32", "Const", {}, {}, &want); + AddNode("y32", "Const", {}, {}, &want); + AddNode("x64", "Const", {}, {}, &want); + AddNode("y64", "Const", {}, {}, &want); + AddNode("out32", "Pow", {"x32", "y32"}, {}, &want); + AddNode("out64", "Pow", {"x64", "y64"}, {}, &want); + + CompareGraphs(want, got); +} + TEST_F(ArithmeticOptimizerTest, Log1p) { tensorflow::Scope s = tensorflow::Scope::NewRootScope(); diff --git a/tensorflow/python/kernel_tests/math_ops/cwise_ops_binary_test.py b/tensorflow/python/kernel_tests/math_ops/cwise_ops_binary_test.py index da827939a7ac5b..21fdc1db607cab 100644 --- a/tensorflow/python/kernel_tests/math_ops/cwise_ops_binary_test.py +++ b/tensorflow/python/kernel_tests/math_ops/cwise_ops_binary_test.py @@ -898,6 +898,26 @@ def testPowNegativeExponentCpu(self): y = -3 self.evaluate(math_ops.pow(x, y)) + # A scalar -1 exponent must not be rewritten to Reciprocal by Grappler. + with test_util.force_cpu(): + with self.assertRaisesRegex( + errors_impl.InvalidArgumentError, + "Integers to negative integer powers are not allowed", + ): + x = np.array([-1, 1]).astype(dtype) + y = np.array(-1).astype(dtype) + self.evaluate(math_ops.pow(x, y)) + + # A uniform vector -1 exponent must not be rewritten either. + with test_util.force_cpu(): + with self.assertRaisesRegex( + errors_impl.InvalidArgumentError, + "Integers to negative integer powers are not allowed", + ): + x = np.array([-1, 1]).astype(dtype) + y = np.array([-1, -1]).astype(dtype) + self.evaluate(math_ops.pow(x, y)) + def testPowNegativeExponentGpu(self): if not test_util.is_gpu_available(): self.skipTest("Requires GPU") From da199e6a14a2999243101debcaea5ad6b2a0470b Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Tue, 29 Sep 2026 21:57:59 -0700 Subject: [PATCH 18/21] Automated Code Change PiperOrigin-RevId: 990730579 --- .../python/ifrt/ir/transforms/ifrt_compile_atom_program_pass.cc | 1 - 1 file changed, 1 deletion(-) diff --git a/third_party/xla/xla/python/ifrt/ir/transforms/ifrt_compile_atom_program_pass.cc b/third_party/xla/xla/python/ifrt/ir/transforms/ifrt_compile_atom_program_pass.cc index 65d63b56a982d5..304ce658ee7e1e 100644 --- a/third_party/xla/xla/python/ifrt/ir/transforms/ifrt_compile_atom_program_pass.cc +++ b/third_party/xla/xla/python/ifrt/ir/transforms/ifrt_compile_atom_program_pass.cc @@ -15,7 +15,6 @@ limitations under the License. #include #include -#include #include #include #include From 1154e44a086bc7cc096a8d99c44ddd501dda44ae Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Tom=C3=A1s=20Longeri?= Date: Tue, 29 Sep 2026 22:12:27 -0700 Subject: [PATCH 19/21] [Mosaic] Fix tpu.memref_squeeze verification and layout propagation for unusual tilings It previously made unverified assumptions about tiling (that it is of the form (A, 128)(...)) and simply propagated the tiling. This is clearly incorrect. For example, even for tilings of this form, it can result in an invalid layout if a 2D shape is squeezed into 1D and the tile remains 2D PiperOrigin-RevId: 990737389 --- .../xla/xla/mosaic/dialect/tpu/tpu_dialect.cc | 4 +- .../xla/xla/mosaic/dialect/tpu/tpu_ops.cc | 144 +++++++++++------- .../xla/xla/mosaic/dialect/tpu/tpu_ops.td | 5 +- .../xla/xla/mosaic/dialect/tpu/util.cc | 8 +- third_party/xla/xla/mosaic/dialect/tpu/util.h | 2 +- 5 files changed, 97 insertions(+), 66 deletions(-) diff --git a/third_party/xla/xla/mosaic/dialect/tpu/tpu_dialect.cc b/third_party/xla/xla/mosaic/dialect/tpu/tpu_dialect.cc index 85c6553bccec52..c05d6aa3130793 100644 --- a/third_party/xla/xla/mosaic/dialect/tpu/tpu_dialect.cc +++ b/third_party/xla/xla/mosaic/dialect/tpu/tpu_dialect.cc @@ -194,11 +194,11 @@ struct MemRefDimOfSqueeze : public OpRewritePattern { } MemRefType source_type = squeeze_op.getInput().getType(); FAILUREOR_ASSIGN_OR_RETURN( - SmallVector squeezed, + SmallVector squeezed, computeSqueezedDimsChecked(squeeze_op, source_type.getShape(), result_type.getShape())); int64_t source_dim = dim; - for (int squeezed_dim : squeezed) { + for (int64_t squeezed_dim : squeezed) { if (squeezed_dim <= source_dim) { ++source_dim; } diff --git a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc index 8688155bcadb36..1a195e74d5cc6d 100644 --- a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc +++ b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.cc @@ -29,6 +29,7 @@ limitations under the License. #include "llvm/ADT/DenseSet.h" #include "llvm/ADT/FloatingPointMode.h" #include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/Sequence.h" #include "llvm/ADT/SmallVector.h" #include "llvm/Support/Casting.h" #include "mlir/Dialect/Arith/IR/Arith.h" @@ -45,6 +46,7 @@ limitations under the License. #include "mlir/IR/BuiltinTypes.h" #include "mlir/IR/Diagnostics.h" #include "mlir/IR/IRMapping.h" +#include "mlir/IR/MLIRContext.h" #include "mlir/IR/Matchers.h" #include "mlir/IR/OpDefinition.h" #include "mlir/IR/OperationSupport.h" @@ -486,87 +488,113 @@ void MemRefSliceOp::getCanonicalizationPatterns(RewritePatternSet& results, } LogicalResult MemRefSqueezeOp::verify() { - MemRefType source_type = getInput().getType(); - MemRefType target_type = getType(); + MemRefType input_type = getInput().getType(); + MemRefType result_type = getType(); - if (target_type.getMemorySpace() != source_type.getMemorySpace()) { + if (result_type.getMemorySpace() != input_type.getMemorySpace()) { return emitOpError("Memory spaces do not match."); } - if (target_type.getElementType() != source_type.getElementType()) { + if (result_type.getElementType() != input_type.getElementType()) { return emitOpError("Element types don't match."); } - auto source_shape = source_type.getShape(); - auto target_shape = target_type.getShape(); + const ArrayRef input_shape = input_type.getShape(); + const ArrayRef result_shape = result_type.getShape(); + // NOTE: In cases where there is flexibility on which dimension to squeeze, + // such as 1x1x2 -> 1x2 (can squeeze dimension 1 or 2), this may choose + // dimensions that have padding and fail in inferResultLayout. + // TODO(tlongeri): Unify logic in inferMemRefReshape such that we don't + // need computeSqueezedDimsChecked at all. FAILUREOR_ASSIGN_OR_RETURN( - auto squeezed, - computeSqueezedDimsChecked(*this, source_shape, target_shape)); - if (squeezed.empty() && source_shape != target_shape) { + const SmallVector squeezed, + computeSqueezedDimsChecked(*this, input_shape, result_shape)); + if (squeezed.empty() && input_shape != result_shape) { return emitOpError( "Source and target shapes must be the same if no dimensions are " "squeezed."); } - auto source_layout = source_type.getLayout(); - auto target_layout = target_type.getLayout(); - bool has_tiled_layout = isa(source_layout); - if (has_tiled_layout != isa(target_layout)) { - return emitOpError( - "Either both src and dst or none of them should have a tiled layout"); - } - if (has_tiled_layout) { - return verifyTiling(); + FAILUREOR_ASSIGN_OR_RETURN( + const MemRefLayoutAttrInterface expected_result_layout, + inferResultLayout(input_type.getLayout(), squeezed, + [&]() { return emitOpError(); })); + if (result_type.getLayout() != expected_result_layout) { + return emitOpError("Expected result layout to be ") + << expected_result_layout; } return success(); } -mlir::InFlightDiagnostic MemRefSqueezeOp::verifyTiling() { - MemRefType source_type = getInput().getType(); - auto source_shape = source_type.getShape(); - auto target_shape = getType().getShape(); - auto squeezed_or = - computeSqueezedDimsChecked(*this, source_shape, target_shape); - if (failed(squeezed_or)) { - return {}; - } - auto& squeezed = squeezed_or.value(); +FailureOr MemRefSqueezeOp::inferResultLayout( + const MemRefLayoutAttrInterface input_layout, + const ArrayRef squeezed, + const function_ref emit_error) { + // TODO(tlongeri): Make MemRefReshapeOp::inferResultLayout more general and + // use that instead. + MLIRContext* const ctx = input_layout.getContext(); + if (auto tiled_layout = dyn_cast(input_layout)) { + const int64_t input_rank = tiled_layout.getRank(); + const int64_t result_rank = input_rank - squeezed.size(); + + SmallVector tile_strides; + tile_strides.reserve(result_rank); + for (int64_t i = 0; i < input_rank; ++i) { + if (!llvm::is_contained(squeezed, i)) { + tile_strides.push_back(tiled_layout.getTileStrides()[i]); + } + } - auto tiles = cast(source_type.getLayout()).getTiles(); - switch (tiles.size()) { - case 0: - break; - case 1: { - auto tile = tiles.front(); - auto tile_dims = tile.dimensions(); - int first_tiled = source_shape.size() - tile_dims.size(); - for (int dim : squeezed) { - if (dim >= first_tiled) { - int tile_idx = dim - first_tiled; - if (tile_idx < 0 || tile_idx >= static_cast(tile_dims.size())) { - return emitOpError() << "Internal error: tile index out of bounds."; - } - if (tile_dims[tile_idx] != 1) { - return emitOpError() - << "All tiled squeezed dimensions must be of size 1."; + const ArrayRef tiles = tiled_layout.getTiles(); + + if (tiles.size() == 1 && tiles[0].dimensions().size() == 2 && + tiles[0].dimension(0) == 1 && + !llvm::is_contained(squeezed, input_rank - 1) && + llvm::is_contained(squeezed, input_rank - 2) && result_rank >= 2) { + // For legacy reasons, for T(1, B) that squeezes the 2nd minor, infer + // T(1, B) instead of T(B) like in the code below. + return MemRefLayoutAttrInterface(TiledLayoutAttr::get( + ctx, {xla::Tile({1, tiles[0].dimension(1)})}, tile_strides)); + } + + SmallVector result_tiles; + result_tiles.reserve(tiles.size()); + // Perform expansion as in TiledLayoutAttr::getExpandedShape, maintaining a + // mapping from expanded shape dimension to corresponding input dimension. + SmallVector expanded_dims = llvm::to_vector( + llvm::iota_range(0, input_rank, /*Inclusive=*/false)); + for (const xla::Tile& input_tile : tiles) { + const int64_t tile_rank = input_tile.dimensions().size(); + const int64_t expanded_rank = expanded_dims.size(); + SmallVector tile; + for (int64_t i = 0; i < tile_rank; ++i) { + const int64_t input_dim = expanded_dims[expanded_rank - tile_rank + i]; + if (llvm::is_contained(squeezed, input_dim)) { + if (input_tile.dimension(i) != 1) { + return emit_error() << "Dimension " << input_dim + << " is padded but is squeezed."; } + } else { + tile.push_back(input_tile.dimension(i)); } + expanded_dims.push_back(input_dim); } - break; - } - default: { - auto first_tile = tiles.front(); - for (int dim : squeezed) { - int first_tiled = source_shape.size() - first_tile.dimensions().size(); - if (dim >= first_tiled) { - return emitOpError() << "When multiple tiles are present, no tiled " - "dimensions can be squeezed."; - } + if (!tile.empty()) { + result_tiles.emplace_back(tile); } } + return MemRefLayoutAttrInterface( + TiledLayoutAttr::get(ctx, result_tiles, tile_strides)); } - - return {}; + if (auto affine_map_attr = dyn_cast(input_layout); + affine_map_attr && affine_map_attr.isIdentity()) { + const int64_t input_rank = affine_map_attr.getValue().getNumInputs(); + const int64_t result_rank = input_rank - squeezed.size(); + // Untiled squeeze + return MemRefLayoutAttrInterface(AffineMapAttr::get( + AffineMap::getMultiDimIdentityMap(result_rank, ctx))); + } + return emit_error() << "Only tiled or identity layouts supported."; } // Rewrites @@ -608,7 +636,7 @@ struct MemRefSqueezeFoldCast : public OpRewritePattern { MemRefType result_type = op.getType(); FAILUREOR_ASSIGN_OR_RETURN( - SmallVector squeezed, + SmallVector squeezed, computeSqueezedDimsChecked(op, cast_result_type.getShape(), result_type.getShape())); diff --git a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td index e31e411588e5c3..1ed0705e1fdde6 100644 --- a/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td +++ b/third_party/xla/xla/mosaic/dialect/tpu/tpu_ops.td @@ -1392,7 +1392,10 @@ def TPU_MemRefSqueezeOp : TPU_Op<"memref_squeeze", [Pure]> { $input attr-dict `:` type($input) `->` type($result) }]; let extraClassDeclaration = [{ - mlir::InFlightDiagnostic verifyTiling(); + static ::mlir::FailureOr<::mlir::MemRefLayoutAttrInterface> inferResultLayout( + ::mlir::MemRefLayoutAttrInterface source_layout, + ::mlir::ArrayRef squeezed, + ::mlir::function_ref<::mlir::InFlightDiagnostic()> emit_error); }]; let hasVerifier = 1; let hasCanonicalizer = 1; diff --git a/third_party/xla/xla/mosaic/dialect/tpu/util.cc b/third_party/xla/xla/mosaic/dialect/tpu/util.cc index 209de7fbd05a1a..600780f81352e0 100644 --- a/third_party/xla/xla/mosaic/dialect/tpu/util.cc +++ b/third_party/xla/xla/mosaic/dialect/tpu/util.cc @@ -35,12 +35,12 @@ limitations under the License. namespace mlir::tpu { -FailureOr> computeSqueezedDimsChecked( +FailureOr> computeSqueezedDimsChecked( Operation* op, ArrayRef source_shape, ArrayRef target_shape) { - SmallVector squeezed; - int source_index = source_shape.size() - 1; - int target_index = target_shape.size() - 1; + SmallVector squeezed; + int64_t source_index = source_shape.size() - 1; + int64_t target_index = target_shape.size() - 1; while (source_index >= 0 || target_index >= 0) { int64_t target_dim = (target_index >= 0) ? target_shape[target_index] : -1; diff --git a/third_party/xla/xla/mosaic/dialect/tpu/util.h b/third_party/xla/xla/mosaic/dialect/tpu/util.h index 18aab9e1b65019..ec3896f5017617 100644 --- a/third_party/xla/xla/mosaic/dialect/tpu/util.h +++ b/third_party/xla/xla/mosaic/dialect/tpu/util.h @@ -167,7 +167,7 @@ std::string shapeToString(const T& shape) { // Computes the dimensions that were squeezed from the source shape to match the // target shape. Returns the dimensions in increasing order. -FailureOr> computeSqueezedDimsChecked( +FailureOr> computeSqueezedDimsChecked( Operation* op, ArrayRef source_shape, ArrayRef target_shape); From 10141aa22d69c6b59f35e4f04b575685a7477992 Mon Sep 17 00:00:00 2001 From: Dirk Hornung Date: Tue, 29 Sep 2026 22:17:53 -0700 Subject: [PATCH 20/21] Preserve metadata and frontend attributes in AllGatherDecomposer. PiperOrigin-RevId: 990739870 --- .../collectives/all_gather_decomposer.cc | 29 ++++++++---- .../collectives/all_gather_decomposer_test.cc | 46 +++++++++++++++++++ 2 files changed, 66 insertions(+), 9 deletions(-) diff --git a/third_party/xla/xla/hlo/transforms/collectives/all_gather_decomposer.cc b/third_party/xla/xla/hlo/transforms/collectives/all_gather_decomposer.cc index 6b67ddc6dbe596..d0c7490ba80e7f 100644 --- a/third_party/xla/xla/hlo/transforms/collectives/all_gather_decomposer.cc +++ b/third_party/xla/xla/hlo/transforms/collectives/all_gather_decomposer.cc @@ -74,12 +74,14 @@ HloInstruction* AllGatherDecomposer::TranslateAllGatherToAllReducePerOperand( auto dus = comp->AddInstruction(HloInstruction::CreateDynamicUpdateSlice( zero->shape(), zero, operand, start_indices)); - auto ar = comp->AddInstruction(HloInstruction::CreateAllReduce( - dus->shape(), {dus}, - MakeBinaryAdd(dus->shape().element_type(), comp->parent()), - ag.device_list(), - /*constrain_layout=*/ag.constrain_layout(), ag.channel_id(), - ag.use_global_device_ids())); + auto ar = comp->AddInstruction( + HloInstruction::CreateAllReduce( + dus->shape(), {dus}, + MakeBinaryAdd(dus->shape().element_type(), comp->parent()), + ag.device_list(), + /*constrain_layout=*/ag.constrain_layout(), ag.channel_id(), + ag.use_global_device_ids()), + &ag.metadata(), &ag.frontend_attributes()); return ar; } @@ -99,14 +101,23 @@ absl::Status AllGatherDecomposer::DecomposeAllGather( tuple_inputs.push_back(ar); } auto tup = comp->AddInstruction(HloInstruction::CreateTuple(tuple_inputs)); - ABSL_RETURN_IF_ERROR(ag->ReplaceAllUsesWith(tup)); + ABSL_RETURN_IF_ERROR( + comp->ReplaceInstruction(ag, tup, /*preserve_sharding=*/false, + /*relay_control_dependency=*/true, + /*remove_unused_operands=*/true, + /*preserve_frontend_attributes=*/false) + .status()); } else { auto* ar = TranslateAllGatherToAllReducePerOperand( group_mode, *ag, ag->shape(), ag->mutable_operand(0), comp, ag->all_gather_dimension()); - ABSL_RETURN_IF_ERROR(ag->ReplaceAllUsesWith(ar)); + ABSL_RETURN_IF_ERROR( + comp->ReplaceInstruction(ag, ar, /*preserve_sharding=*/false, + /*relay_control_dependency=*/true, + /*remove_unused_operands=*/true, + /*preserve_frontend_attributes=*/false) + .status()); } - ABSL_RETURN_IF_ERROR(comp->RemoveInstructionAndUnusedOperands(ag)); return absl::OkStatus(); } diff --git a/third_party/xla/xla/hlo/transforms/collectives/all_gather_decomposer_test.cc b/third_party/xla/xla/hlo/transforms/collectives/all_gather_decomposer_test.cc index ca27c1970a04b5..9c24dc0f64f3df 100644 --- a/third_party/xla/xla/hlo/transforms/collectives/all_gather_decomposer_test.cc +++ b/third_party/xla/xla/hlo/transforms/collectives/all_gather_decomposer_test.cc @@ -182,5 +182,51 @@ ENTRY entry { op::Multiply(op::ReplicaId(), op::Constant()))))); } +TEST_F(AllGatherDecomposerTest, PreservesMetadataAndFrontendAttributes) { + const std::string module_str = R"( +HloModule module + +ENTRY entry { + param0 = f32[10,20] parameter(0) + param1 = f32[10,16] parameter(1) + ag_single = f32[10,80] all-gather(param0), replica_groups={}, dimensions={1}, + frontend_attributes={_xla_compute_type="sparse"}, + metadata={op_name="jit(all_gather_fn)/shard_map/MARKER!!!/all_gather"} + ag_tuple = (f32[10,80], f32[10,64]) all-gather(param0, param1), + replica_groups={}, dimensions={1}, + frontend_attributes={_xla_compute_type="sparse"}, + metadata={op_name="jit(all_gather_fn)/shard_map/MARKER_TUPLE/all_gather"} + ROOT root = (f32[10,80], (f32[10,80], f32[10,64])) tuple(ag_single, ag_tuple) +} +)"; + + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnUnverifiedModule(module_str)); + AllGatherDecomposer decomposer; + ASSERT_OK_AND_ASSIGN(bool changed, decomposer.Run(module.get())); + EXPECT_TRUE(changed); + + const HloInstruction* root = module->entry_computation()->root_instruction(); + const HloInstruction* ar_single = root->operand(0); + EXPECT_THAT(ar_single, op::AllReduce()); + EXPECT_EQ(ar_single->metadata().op_name(), + "jit(all_gather_fn)/shard_map/MARKER!!!/all_gather"); + EXPECT_THAT( + ar_single->frontend_attributes().map(), + ::testing::Contains(::testing::Pair("_xla_compute_type", "sparse"))); + + const HloInstruction* ar_tuple = root->operand(1); + EXPECT_THAT(ar_tuple, op::Tuple(op::AllReduce(), op::AllReduce())); + EXPECT_EQ(ar_tuple->metadata().op_name(), + "jit(all_gather_fn)/shard_map/MARKER_TUPLE/all_gather"); + for (const HloInstruction* ar : ar_tuple->operands()) { + EXPECT_EQ(ar->metadata().op_name(), + "jit(all_gather_fn)/shard_map/MARKER_TUPLE/all_gather"); + EXPECT_THAT( + ar->frontend_attributes().map(), + ::testing::Contains(::testing::Pair("_xla_compute_type", "sparse"))); + } +} + } // namespace } // namespace xla From 103b097f3228df939435267c6f90481c287833fd Mon Sep 17 00:00:00 2001 From: Andrew Dame Date: Tue, 29 Sep 2026 23:07:36 -0700 Subject: [PATCH 21/21] Allow concurrent `workflow_dispatch` runs in `postsubmit_benchmark.yml` Updates the concurrency group to append the run ID when the workflow is triggered via workflow_dispatch. This allows manual runs (such as autobisect) to execute in parallel, while push events continue to run sequentially to prevent clashes. PiperOrigin-RevId: 990760148 --- third_party/xla/.github/workflows/postsubmit_benchmark.yml | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/third_party/xla/.github/workflows/postsubmit_benchmark.yml b/third_party/xla/.github/workflows/postsubmit_benchmark.yml index 71a35554a5fecc..7b577eca070997 100644 --- a/third_party/xla/.github/workflows/postsubmit_benchmark.yml +++ b/third_party/xla/.github/workflows/postsubmit_benchmark.yml @@ -50,10 +50,10 @@ permissions: pull-requests: write concurrency: - # Group by workflow name and branch to ensure runs execute sequentially. + # Group by workflow name and branch on push to ensure runs execute sequentially. # This prevents parallel postsubmit runs from clashing and ensures accurate - # reporting history. - group: ${{ github.workflow }}-${{ github.ref }} + # reporting history. For workflow_dispatch (autobisect), append the run_id so they run in parallel without cancelling each other. + group: ${{ github.workflow }}-${{ github.ref }}-${{ github.event_name == 'workflow_dispatch' && github.run_id || 'push' }} cancel-in-progress: false jobs: