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/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( 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/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 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 { 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 b1f041d562fe51..a4d63b3b4c0636 100644 --- a/tensorflow/core/kernels/mlir_generated/gpu_binary_ops_test.cc +++ b/tensorflow/core/kernels/mlir_generated/gpu_binary_ops_test.cc @@ -1503,7 +1503,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, @@ -1526,12 +1526,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, @@ -1576,7 +1577,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, 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/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) { 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; 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") diff --git a/tensorflow/python/ops/math_ops_test.py b/tensorflow/python/ops/math_ops_test.py index 1fbc1795a0b986..734994becd57f2 100644 --- a/tensorflow/python/ops/math_ops_test.py +++ b/tensorflow/python/ops/math_ops_test.py @@ -1104,6 +1104,51 @@ 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.0), + (math_ops.xlog1py, 1.0), + (math_ops.xdivy, 3.0), + ): + for dtype in [dtypes.float16, dtypes.float32, dtypes.float64]: + x = constant_op.constant(np.full((16,), -0.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.0, 1.0]).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, 0.0], dtype=dtype) + y = constant_op.constant([1.0, -1.0], dtype=dtype) + with test_util.force_gpu(): + result = self.evaluate(math_ops.xlog1py(x, y)) + self.assertAllClose(result, [np.log(2.0), 0.0]) + + @test_util.run_all_in_graph_and_eager_modes class XlogyTest(test_util.TensorFlowTestCase): 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: diff --git a/third_party/xla/third_party/stablehlo/temporary.patch b/third_party/xla/third_party/stablehlo/temporary.patch index 73906a9c97345e..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 @@ -548,6 +570,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 @@ -822,6 +878,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 @@ -937,10 +1004,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 ), [{ @@ -952,6 +1080,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 +1104,22 @@ 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,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; + 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 @@ -975,6 +1131,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 +1196,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 @@ -1138,6 +1343,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 @@ -2132,6 +2450,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 +2556,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/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/backends/cpu/codegen/ir_compiler.cc b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc index b99fbcce5cbe15..c59d2cded0e085 100644 --- a/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc +++ b/third_party/xla/xla/backends/cpu/codegen/ir_compiler.cc @@ -472,13 +472,11 @@ 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); + // 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/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); 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/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_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 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/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/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); 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 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/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/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> 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 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 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); 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;