From a55459e27e7e7fac3de8be8f1325259109a3cbe6 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Thu, 23 Jul 2026 11:40:30 -0700 Subject: [PATCH 01/58] Reject malformed empty boxes and box_index in crop-and-resize kernels ParseAndCheckBoxSizes returned early with num_boxes=0 when both boxes and box_index were empty, skipping all rank validation. A rank-1 empty boxes tensor (shape [0] instead of [0, 4]) or a rank-2 empty box_index then reached Tensor::tensor() with the wrong rank, aborting the process with a fatal CHECK failure instead of raising a catchable InvalidArgumentError. This crashed CropAndResizeGradImage and CropAndResizeGradBoxes. Validate the ranks of boxes and box_index before the empty early return. Well-formed empty inputs (boxes of shape [0, 4] with box_index of shape [0]) keep working, and empty tensors of the correct rank are still accepted. Fixes #123397 --- .../core/kernels/image/crop_and_resize_op.cc | 21 +++++---- tensorflow/python/ops/image_grad_test_base.py | 46 +++++++++++++++++++ 2 files changed, 57 insertions(+), 10 deletions(-) diff --git a/tensorflow/core/kernels/image/crop_and_resize_op.cc b/tensorflow/core/kernels/image/crop_and_resize_op.cc index ff09def018cddb..a01432d96de803 100644 --- a/tensorflow/core/kernels/image/crop_and_resize_op.cc +++ b/tensorflow/core/kernels/image/crop_and_resize_op.cc @@ -55,24 +55,25 @@ using Callback = std::function; static inline absl::Status ParseAndCheckBoxSizes(const Tensor& boxes, const Tensor& box_index, int* num_boxes) { - if (boxes.NumElements() == 0 && box_index.NumElements() == 0) { - *num_boxes = 0; - return absl::OkStatus(); - } - // The shape of 'boxes' is [num_boxes, 4]. + // The shape of 'boxes' is [num_boxes, 4] and the shape of 'box_index' is + // [num_boxes]. The ranks must be validated even when both tensors are + // empty, since the kernels later access them as rank-2 and rank-1 tensors. if (boxes.dims() != 2) { return absl::InvalidArgumentError( absl::StrCat("boxes must be 2-D", boxes.shape().DebugString())); } - *num_boxes = boxes.dim_size(0); - if (boxes.dim_size(1) != 4) { - return absl::InvalidArgumentError("boxes must have 4 columns"); - } - // The shape of 'box_index' is [num_boxes]. if (box_index.dims() != 1) { return absl::InvalidArgumentError( absl::StrCat("box_index must be 1-D", box_index.shape().DebugString())); } + if (boxes.NumElements() == 0 && box_index.NumElements() == 0) { + *num_boxes = 0; + return absl::OkStatus(); + } + *num_boxes = boxes.dim_size(0); + if (boxes.dim_size(1) != 4) { + return absl::InvalidArgumentError("boxes must have 4 columns"); + } if (box_index.dim_size(0) != *num_boxes) { return absl::InvalidArgumentError("box_index has incompatible shape"); } diff --git a/tensorflow/python/ops/image_grad_test_base.py b/tensorflow/python/ops/image_grad_test_base.py index fc84f3054f9656..65f5c7f157fb0e 100644 --- a/tensorflow/python/ops/image_grad_test_base.py +++ b/tensorflow/python/ops/image_grad_test_base.py @@ -24,6 +24,7 @@ from tensorflow.python.framework import dtypes from tensorflow.python.framework import errors_impl from tensorflow.python.framework import test_util +from tensorflow.python.ops import array_ops from tensorflow.python.ops import array_ops_stack from tensorflow.python.ops import gen_image_ops from tensorflow.python.ops import gradient_checker_v2 @@ -447,6 +448,51 @@ def testShapeIsCorrectAfterOp(self): crops = self.evaluate(crops) self.assertEqual(crops_shape, list(crops.shape)) + def testMalformedEmptyBoxesRaisesError(self): + # Regression test for GitHub issue 123397: rank-1 empty `boxes` or + # rank-2 empty `box_ind` used to pass validation and crash the process + # with a fatal CHECK failure instead of raising InvalidArgumentError. + grads = array_ops.zeros([0, 1, 1, 1], dtype=dtypes.float32) + image_size = constant_op.constant([2, 7, 7, 1], dtype=dtypes.int32) + valid_boxes = array_ops.zeros([0, 4], dtype=dtypes.float32) + valid_box_ind = array_ops.zeros([0], dtype=dtypes.int32) + with self.assertRaisesRegex( + (errors_impl.InvalidArgumentError, ValueError), "boxes must be 2-D"): + self.evaluate( + gen_image_ops.crop_and_resize_grad_image( + grads=grads, + boxes=array_ops.zeros([0], dtype=dtypes.float32), + box_ind=valid_box_ind, + image_size=image_size, + T=dtypes.float32)) + with self.assertRaisesRegex( + (errors_impl.InvalidArgumentError, ValueError), "box_index must be 1-D" + ): + self.evaluate( + gen_image_ops.crop_and_resize_grad_image( + grads=grads, + boxes=valid_boxes, + box_ind=array_ops.zeros([0, 0], dtype=dtypes.int32), + image_size=image_size, + T=dtypes.float32)) + with self.assertRaisesRegex( + (errors_impl.InvalidArgumentError, ValueError), "boxes must be 2-D"): + self.evaluate( + gen_image_ops.crop_and_resize_grad_boxes( + grads=grads, + image=array_ops.zeros([2, 7, 7, 1], dtype=dtypes.float32), + boxes=array_ops.zeros([0], dtype=dtypes.float32), + box_ind=valid_box_ind)) + # Well-formed empty boxes must keep working. + output = self.evaluate( + gen_image_ops.crop_and_resize_grad_image( + grads=grads, + boxes=valid_boxes, + box_ind=valid_box_ind, + image_size=image_size, + T=dtypes.float32)) + self.assertEqual((2, 7, 7, 1), output.shape) + def _randomUniformAvoidAnchors(self, low, high, anchors, radius, num_samples): """Generate samples that are far enough from a set of anchor points. From b7d372e6490129ad8fa1bdcc2c487b52bf45d053 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Thu, 23 Jul 2026 11:47:32 -0700 Subject: [PATCH 02/58] Add separator before shape in boxes and box_index rank error messages The messages previously rendered as "boxes must be 2-D[0]". They now read "boxes must be 2-D, got [0]". --- tensorflow/core/kernels/image/crop_and_resize_op.cc | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tensorflow/core/kernels/image/crop_and_resize_op.cc b/tensorflow/core/kernels/image/crop_and_resize_op.cc index a01432d96de803..978a1ab10b8358 100644 --- a/tensorflow/core/kernels/image/crop_and_resize_op.cc +++ b/tensorflow/core/kernels/image/crop_and_resize_op.cc @@ -59,12 +59,12 @@ static inline absl::Status ParseAndCheckBoxSizes(const Tensor& boxes, // [num_boxes]. The ranks must be validated even when both tensors are // empty, since the kernels later access them as rank-2 and rank-1 tensors. if (boxes.dims() != 2) { - return absl::InvalidArgumentError( - absl::StrCat("boxes must be 2-D", boxes.shape().DebugString())); + return absl::InvalidArgumentError(absl::StrCat( + "boxes must be 2-D, got ", boxes.shape().DebugString())); } if (box_index.dims() != 1) { - return absl::InvalidArgumentError( - absl::StrCat("box_index must be 1-D", box_index.shape().DebugString())); + return absl::InvalidArgumentError(absl::StrCat( + "box_index must be 1-D, got ", box_index.shape().DebugString())); } if (boxes.NumElements() == 0 && box_index.NumElements() == 0) { *num_boxes = 0; From 4e6673452fbcf430b6567f3b72d21803d49df74f Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Sat, 22 Aug 2026 18:37:10 -0700 Subject: [PATCH 03/58] Add missing strict dep for image_grad_test_base --- tensorflow/python/ops/BUILD | 1 + 1 file changed, 1 insertion(+) diff --git a/tensorflow/python/ops/BUILD b/tensorflow/python/ops/BUILD index db5e44573e8878..e5f728aa37b1ef 100644 --- a/tensorflow/python/ops/BUILD +++ b/tensorflow/python/ops/BUILD @@ -3486,6 +3486,7 @@ py_library( srcs = ["image_grad_test_base.py"], strict_deps = True, deps = [ + ":array_ops", ":array_ops_stack", ":gradient_checker_v2", ":image_ops", From 77637cc026338d00539b6e90147665014fdf6e61 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Wed, 26 Aug 2026 00:08:10 -0700 Subject: [PATCH 04/58] Reject SplitV size_splits overflow instead of aborting --- tensorflow/core/kernels/split_v_op.cc | 16 ++++++++++++++ .../python/kernel_tests/array_ops/BUILD | 1 + .../kernel_tests/array_ops/split_op_test.py | 21 +++++++++++++++++++ 3 files changed, 38 insertions(+) diff --git a/tensorflow/core/kernels/split_v_op.cc b/tensorflow/core/kernels/split_v_op.cc index bd89be3b2dff37..640963f981bfd2 100644 --- a/tensorflow/core/kernels/split_v_op.cc +++ b/tensorflow/core/kernels/split_v_op.cc @@ -27,6 +27,7 @@ limitations under the License. #define PLUGGABLE_DEVICE_SUPPORTED_MACOS 1 #endif +#include #include #include "unsupported/Eigen/CXX11/Tensor" // from @eigen_archive @@ -128,6 +129,21 @@ class SplitVOpBase : public OpKernel { "input.")); neg_one_dim = d; } else { + // Accumulate with an explicit overflow guard. `determined_size += size` + // wraps for large `size_splits`, and a wrapped total can equal + // `input_size_split_dim` and pass the check below, letting the aligned + // slicing path compute an endpoint that reaches a fatal `Tensor::Slice` + // invariant. Rejecting overflow here also keeps that later path safe, + // since the accepted total then bounds every partial sum. + const bool overflow = + (size > 0 && + determined_size > std::numeric_limits::max() - size) || + (size < 0 && + determined_size < std::numeric_limits::min() - size); + OP_REQUIRES(context, !overflow, + absl::InvalidArgumentError(absl::StrCat( + "Sum of size_splits overflows the index type at index ", + d, "."))); determined_size += size; } } diff --git a/tensorflow/python/kernel_tests/array_ops/BUILD b/tensorflow/python/kernel_tests/array_ops/BUILD index 72f6e96c5a8434..3e1a2513354bd2 100644 --- a/tensorflow/python/kernel_tests/array_ops/BUILD +++ b/tensorflow/python/kernel_tests/array_ops/BUILD @@ -767,6 +767,7 @@ cuda_py_strict_test( "//tensorflow/python/framework:for_generated_wrappers", "//tensorflow/python/framework:test_lib", "//tensorflow/python/ops:array_ops", + "//tensorflow/python/ops:array_ops_gen", "//tensorflow/python/ops:gradients_impl", "//tensorflow/python/ops:math_ops", "//tensorflow/python/platform:client_testlib", diff --git a/tensorflow/python/kernel_tests/array_ops/split_op_test.py b/tensorflow/python/kernel_tests/array_ops/split_op_test.py index 9c2b9d2ac6723d..89464c213a87a8 100644 --- a/tensorflow/python/kernel_tests/array_ops/split_op_test.py +++ b/tensorflow/python/kernel_tests/array_ops/split_op_test.py @@ -22,6 +22,7 @@ from tensorflow.python.framework import ops from tensorflow.python.framework import test_util from tensorflow.python.ops import array_ops +from tensorflow.python.ops import gen_array_ops from tensorflow.python.ops import gradients_impl from tensorflow.python.ops import math_ops from tensorflow.python.platform import test @@ -120,6 +121,26 @@ def testExplicitNum(self): self.assertAllEqual(r[1], value[2:4]) self.assertAllEqual(r[2], value[4:]) + @test_util.run_in_graph_and_eager_modes + def testSizeSplitsOverflowRaises(self): + # Regression test for GitHub issue 126126. The cumulative sum of + # size_splits was computed with unchecked signed addition, so a wrapped + # total could equal the input dimension, pass validation, and reach a + # fatal `Tensor::Slice` invariant in the aligned slicing path. It must + # raise instead. + i64_max = (1 << 63) - 1 + value = array_ops.reshape( + constant_op.constant([], dtype=dtypes.float32), + constant_op.constant([i64_max, 0], dtype=dtypes.int64)) + with self.assertRaisesRegex(errors_impl.InvalidArgumentError, "overflow"): + self.evaluate( + gen_array_ops.SplitV( + value=value, + size_splits=constant_op.constant( + [i64_max, i64_max, i64_max, 2], dtype=dtypes.int64), + axis=constant_op.constant(0, dtype=dtypes.int32), + num_split=4)) + @test_util.run_in_graph_and_eager_modes def testListOfScalarTensors(self): a = math_ops.cast(5, dtypes.int32) From 86e8344bf4d34a776dbc00484e7810f4c47d48e9 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Wed, 26 Aug 2026 09:34:07 -0700 Subject: [PATCH 05/58] Skip the SplitV overflow test when XLA is enabled --- tensorflow/python/kernel_tests/array_ops/split_op_test.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/tensorflow/python/kernel_tests/array_ops/split_op_test.py b/tensorflow/python/kernel_tests/array_ops/split_op_test.py index 89464c213a87a8..13b2dea79a76e9 100644 --- a/tensorflow/python/kernel_tests/array_ops/split_op_test.py +++ b/tensorflow/python/kernel_tests/array_ops/split_op_test.py @@ -128,6 +128,10 @@ def testSizeSplitsOverflowRaises(self): # total could equal the input dimension, pass validation, and reach a # fatal `Tensor::Slice` invariant in the aligned slicing path. It must # raise instead. + if test_util.is_xla_enabled(): + # XLA cannot compile the reshape to an INT64_MAX dimension, so the test + # cannot reach the SplitV kernel under XLA. + self.skipTest("XLA does not support shapes whose element count overflows") i64_max = (1 << 63) - 1 value = array_ops.reshape( constant_op.constant([], dtype=dtypes.float32), From b47c91eb743b5f5658b0fd2d7f80605323e5b1a3 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Wed, 26 Aug 2026 09:38:35 -0700 Subject: [PATCH 06/58] Use disable_xla decorator for the SplitV overflow test --- tensorflow/python/kernel_tests/array_ops/split_op_test.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/tensorflow/python/kernel_tests/array_ops/split_op_test.py b/tensorflow/python/kernel_tests/array_ops/split_op_test.py index 13b2dea79a76e9..0449952dd11f61 100644 --- a/tensorflow/python/kernel_tests/array_ops/split_op_test.py +++ b/tensorflow/python/kernel_tests/array_ops/split_op_test.py @@ -122,16 +122,15 @@ def testExplicitNum(self): self.assertAllEqual(r[2], value[4:]) @test_util.run_in_graph_and_eager_modes + @test_util.disable_xla( + "XLA shape inference rejects the reshape to an INT64_MAX dimension, so " + "the test cannot reach the SplitV kernel under XLA") def testSizeSplitsOverflowRaises(self): # Regression test for GitHub issue 126126. The cumulative sum of # size_splits was computed with unchecked signed addition, so a wrapped # total could equal the input dimension, pass validation, and reach a # fatal `Tensor::Slice` invariant in the aligned slicing path. It must # raise instead. - if test_util.is_xla_enabled(): - # XLA cannot compile the reshape to an INT64_MAX dimension, so the test - # cannot reach the SplitV kernel under XLA. - self.skipTest("XLA does not support shapes whose element count overflows") i64_max = (1 << 63) - 1 value = array_ops.reshape( constant_op.constant([], dtype=dtypes.float32), From 2ca8a3a10d07cd251d212cf1439a55f8baf49d4a Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Mon, 21 Sep 2026 11:45:49 -0700 Subject: [PATCH 07/58] Stop exporting statically linked LLVM and MLIR symbols from libtensorflow_framework libtensorflow_framework.so is linked with tf_framework_version_script.lds, which read: tensorflow { global: *; }; That declares every symbol global and has no local stanza, so the library exports its statically linked third-party code as well as its own. Among those are the LLVM symbols, which then reach the process-wide resolution scope and collide with any other LLVM loaded into the same process. In the reported crash an LLVM 15 brought in by llvmlite and numba resolved llvm::raw_svector_ostream::write_impl against TensorFlow's LLVM 18 and died on the mismatched layout. Add a local stanza hiding *llvm* and *mlir*, and add *llvm* to tf_private_symbols.lds so macOS, which already hides *mlir* through -unexported_symbols_list, covers the same set. The change is deliberately narrow. TensorFlow's own symbols stay exported, because custom op libraries loaded through tf.load_op_library resolve REGISTER_OP and REGISTER_KERNEL_BUILDER against this library, which is why the script exports everything today. Verified with lld, the linker the Linux builds use via -fuse-ld=lld, on a shared object built from the shipped script: llvm and mlir symbols are dropped from .dynsym while the tensorflow namespace and the TF_ C API stay exported. Linking also succeeds, with and without --undefined-version, when no symbol matches either pattern, so configurations built without LLVM are unaffected. Fixes #104038 --- tensorflow/tf_framework_version_script.lds | 16 +++++++++++++++- tensorflow/tf_private_symbols.lds | 1 + 2 files changed, 16 insertions(+), 1 deletion(-) diff --git a/tensorflow/tf_framework_version_script.lds b/tensorflow/tf_framework_version_script.lds index 99ed72972e3ba6..e557809002bfa9 100644 --- a/tensorflow/tf_framework_version_script.lds +++ b/tensorflow/tf_framework_version_script.lds @@ -1,4 +1,18 @@ tensorflow { global: *; -}; \ No newline at end of file + + # Statically linked third-party code whose symbols must not reach the + # process-wide resolution scope. LLVM in particular collides with other + # copies loaded into the same process, for example the LLVM that llvmlite + # and numba bring in, which resolve against TensorFlow's build and crash on + # the mismatched layout. + # + # These mirror the entries already hidden on macOS in tf_private_symbols.lds + # and are deliberately narrow: TensorFlow's own symbols stay exported, + # because custom op libraries loaded through tf.load_op_library resolve + # REGISTER_OP and REGISTER_KERNEL_BUILDER against this library. + local: + *llvm*; + *mlir*; +}; diff --git a/tensorflow/tf_private_symbols.lds b/tensorflow/tf_private_symbols.lds index 319b40bc72b66a..572525f7c7d4fe 100644 --- a/tensorflow/tf_private_symbols.lds +++ b/tensorflow/tf_private_symbols.lds @@ -6,4 +6,5 @@ _jzero_far _jcopy_* _jsimd_* _hwloc_* +*llvm* *mlir* From f2c41de51323c4fe66d46190b64ff243de68454c Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Mon, 21 Sep 2026 13:01:26 -0700 Subject: [PATCH 08/58] Match LLVM and MLIR by namespace instead of by substring Per review, the substring patterns were both too broad and too narrow. Too broad: *mlir* also hides TensorFlow's own symbols whose names happen to contain the string, such as tensorflow::tf_xla_test_use_mlir, which would become local and fail to resolve for anything linking against it. Too narrow: *llvm* is lowercase and never matched the LLVM C API at all, whose symbols are spelled LLVM*. Those kept leaking into the dynamic symbol table, which is the same collision class the change is meant to close. Match the namespaces through extern "C++" and add an explicit LLVM* prefix for the C API. Verified with lld on the shipped script: llvm::, mlir:: and LLVM* symbols are dropped while tensorflow::, including tf_xla_test_use_mlir, and the TF_ C API stay exported. Linking still succeeds when no pattern matches. Also adds *LLVM* to the macOS list, which had the same case gap next to its existing *mlir* entry. --- tensorflow/tf_framework_version_script.lds | 14 ++++++++++---- tensorflow/tf_private_symbols.lds | 1 + 2 files changed, 11 insertions(+), 4 deletions(-) diff --git a/tensorflow/tf_framework_version_script.lds b/tensorflow/tf_framework_version_script.lds index e557809002bfa9..9d933db6f770d7 100644 --- a/tensorflow/tf_framework_version_script.lds +++ b/tensorflow/tf_framework_version_script.lds @@ -8,11 +8,17 @@ tensorflow { # and numba bring in, which resolve against TensorFlow's build and crash on # the mismatched layout. # - # These mirror the entries already hidden on macOS in tf_private_symbols.lds - # and are deliberately narrow: TensorFlow's own symbols stay exported, + # Matched by namespace rather than by substring so that TensorFlow's own + # symbols are not caught: a plain *mlir* would also hide names such as + # tensorflow::tf_xla_test_use_mlir. TensorFlow's symbols must stay exported, # because custom op libraries loaded through tf.load_op_library resolve # REGISTER_OP and REGISTER_KERNEL_BUILDER against this library. + # + # LLVM* covers the LLVM C API, whose symbols are not in the llvm namespace. local: - *llvm*; - *mlir*; + extern "C++" { + llvm::*; + mlir::*; + }; + LLVM*; }; diff --git a/tensorflow/tf_private_symbols.lds b/tensorflow/tf_private_symbols.lds index 572525f7c7d4fe..488eed51d27869 100644 --- a/tensorflow/tf_private_symbols.lds +++ b/tensorflow/tf_private_symbols.lds @@ -7,4 +7,5 @@ _jcopy_* _jsimd_* _hwloc_* *llvm* +*LLVM* *mlir* From 1b5868f96fbd9a5402d227e2d2ea951be9fb5125 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Fri, 18 Sep 2026 20:32:40 -0700 Subject: [PATCH 09/58] Validate axes in np.rot90 for bounds, duplicates, and length --- .../python/ops/numpy_ops/np_array_ops.py | 17 ++++++++++++++++ .../python/ops/numpy_ops/tests/np_test.py | 20 +++++++++++++++++++ 2 files changed, 37 insertions(+) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index 250893b1e4797c..5201cfb7b87b18 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -1582,6 +1582,23 @@ def roll(a, shift, axis=None): # pylint: disable=missing-docstring @tf_export.tf_export('experimental.numpy.rot90', v1=[]) @np_utils.np_doc('rot90') def rot90(m, k=1, axes=(0, 1)): # pylint: disable=missing-docstring + m = asarray(m) + + maybe_rank = m.shape.rank + if maybe_rank is not None: + if not isinstance(axes, (tuple, list, range, np.ndarray)): + axes = tuple(axes) + if len(axes) != 2: + raise ValueError('len(axes) must be 2.') + if builtins.all(isinstance(axis, (int, np.integer)) for axis in axes): + norm_axes = [axis + maybe_rank if axis < 0 else axis for axis in axes] + if builtins.any(axis < 0 or axis >= maybe_rank for axis in norm_axes): + raise ValueError( + f'Axes={tuple(axes)} out of range for array of rank {maybe_rank}.' + ) + if norm_axes[0] == norm_axes[1]: + raise ValueError('Axes must be different.') + m_rank = array_ops.rank(m) ax1, ax2 = np_utils._canonicalize_axes(axes, m_rank) # pylint: disable=protected-access diff --git a/tensorflow/python/ops/numpy_ops/tests/np_test.py b/tensorflow/python/ops/numpy_ops/tests/np_test.py index 054549a78fea00..bd6fd6f4f960bf 100644 --- a/tensorflow/python/ops/numpy_ops/tests/np_test.py +++ b/tensorflow/python/ops/numpy_ops/tests/np_test.py @@ -2069,6 +2069,26 @@ def testRot90Additional(self, shape, dtype, k, axes, rng_factory): self._CompileAndCheck( lnp_op, args_maker, check_dtypes=True, check_incomplete_shape=True) + @new_test + def testRot90InvalidAxes(self): + a = tnp.ones((2, 3)) + # Out-of-bounds axes must be rejected, matching NumPy's ValueError. + with self.assertRaisesRegex(ValueError, 'out of range'): + tnp.rot90(a, axes=(0, 3)) + with self.assertRaisesRegex(ValueError, 'out of range'): + tnp.rot90(a, axes=(0, -5)) + # Duplicate axes (after negative normalization) must be rejected. + with self.assertRaisesRegex(ValueError, 'must be different'): + tnp.rot90(a, axes=(0, 0)) + with self.assertRaisesRegex(ValueError, 'must be different'): + tnp.rot90(a, axes=(0, -2)) + # axes must have exactly two entries. + with self.assertRaisesRegex(ValueError, 'must be 2'): + tnp.rot90(a, axes=(0,)) + # In-bounds negative axes remain valid. + self.assertAllClose(tnp.rot90(a, axes=(-2, -1)), onp.rot90( + onp.ones((2, 3)), axes=(-2, -1))) + # TODO(mattjj): test infix operator overrides def testRavel(self): From 1d3b095e8c41f1c3426d018fe30f96bf27e1582e Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Mon, 21 Sep 2026 18:57:14 -0700 Subject: [PATCH 10/58] Skip rot90 axes validation for non-sequence axes --- tensorflow/python/ops/numpy_ops/np_array_ops.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index 5201cfb7b87b18..14a110fb531656 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -1585,9 +1585,9 @@ def rot90(m, k=1, axes=(0, 1)): # pylint: disable=missing-docstring m = asarray(m) maybe_rank = m.shape.rank - if maybe_rank is not None: - if not isinstance(axes, (tuple, list, range, np.ndarray)): - axes = tuple(axes) + if maybe_rank is not None and isinstance( + axes, (tuple, list, range, np.ndarray) + ): if len(axes) != 2: raise ValueError('len(axes) must be 2.') if builtins.all(isinstance(axis, (int, np.integer)) for axis in axes): From c64ede3baf9d8ba7ed193248e2e494ee9cb4b71d Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Tue, 22 Sep 2026 06:41:04 -0700 Subject: [PATCH 11/58] Wrap nonempty_nonscalar_array_shapes to satisfy 80-char limit --- tensorflow/python/ops/numpy_ops/tests/np_test.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tensorflow/python/ops/numpy_ops/tests/np_test.py b/tensorflow/python/ops/numpy_ops/tests/np_test.py index bd6fd6f4f960bf..c06e4f9139651e 100644 --- a/tensorflow/python/ops/numpy_ops/tests/np_test.py +++ b/tensorflow/python/ops/numpy_ops/tests/np_test.py @@ -39,7 +39,9 @@ config.parse_flags_with_absl() -nonempty_nonscalar_array_shapes = [(4,), (3, 4), (3, 1), (1, 4), (2, 1, 4), (2, 3, 4)] +nonempty_nonscalar_array_shapes = [ + (4,), (3, 4), (3, 1), (1, 4), (2, 1, 4), (2, 3, 4) +] nonempty_array_shapes = [()] + nonempty_nonscalar_array_shapes empty_array_shapes = [(0,), (0, 4), (3, 0),] From 779544ae91500a48c136ea62dfd233946a08cb4b Mon Sep 17 00:00:00 2001 From: Rohith Pariki Date: Thu, 24 Sep 2026 04:33:46 +0530 Subject: [PATCH 12/58] Validate filter element count against 32-bit limit in GPU conv transformations Filter transformation kernels on GPU (TransformFilter and ReverseTransformFilter) use 32-bit indexing via To32Bit. When filter element count exceeds INT32_MAX, signed integer overflow caused fatal process crashes in gpu_launch_config.h (CHECK_GE work_element_count >= 0) and out-of-bounds illegal memory access in ShuffleInTensor3Simple. This change adds FastBoundsCheck upfront validation in conv_ops_impl.h, conv_ops_fused_impl.h, conv_grad_input_ops.cc, and conv_grad_input_ops_3d.cc, and guards non-positive output sizes in conv_2d_gpu.h. Fixes #87457 Fixes #87438 Fixes #87454 --- tensorflow/core/kernels/conv_2d_gpu.h | 7 +++++++ tensorflow/core/kernels/conv_grad_input_ops.cc | 6 ++++++ .../core/kernels/conv_grad_input_ops_3d.cc | 9 +++++++++ tensorflow/core/kernels/conv_ops_fused_impl.h | 8 ++++++++ tensorflow/core/kernels/conv_ops_impl.h | 8 ++++++++ .../kernel_tests/nn_ops/conv_ops_test.py | 18 ++++++++++++++++++ 6 files changed, 56 insertions(+) diff --git a/tensorflow/core/kernels/conv_2d_gpu.h b/tensorflow/core/kernels/conv_2d_gpu.h index 3c9ccb5e6330f8..e33294e18546cb 100644 --- a/tensorflow/core/kernels/conv_2d_gpu.h +++ b/tensorflow/core/kernels/conv_2d_gpu.h @@ -491,6 +491,9 @@ struct TransformFilter { } combined_dims[1] = in.dimension(NDIMS - 2); // input filters combined_dims[2] = in.dimension(NDIMS - 1); // output filters + if (TF_PREDICT_FALSE(out.size() <= 0)) { + return; + } GpuLaunchConfig config = GetGpuLaunchConfig(out.size(), d); if (dst_filter_format == FORMAT_OIHW) { @@ -522,6 +525,10 @@ struct ReverseTransformFilter { typename TTypes::Tensor out) { Dimension<3> combined_dims; + if (TF_PREDICT_FALSE(out.size() <= 0)) { + return; + } + if (src_filter_format == FORMAT_OIHW) { combined_dims[0] = in.dimension(0); // output filters combined_dims[1] = in.dimension(1); // input filters diff --git a/tensorflow/core/kernels/conv_grad_input_ops.cc b/tensorflow/core/kernels/conv_grad_input_ops.cc index e56826e7ddf580..93d5cbae5eba23 100644 --- a/tensorflow/core/kernels/conv_grad_input_ops.cc +++ b/tensorflow/core/kernels/conv_grad_input_ops.cc @@ -293,6 +293,12 @@ void LaunchConv2DBackpropInputOpGpuImpl( TF_RETURN_IF_ERROR(ctx->allocate_temp(DataTypeToEnum::value, dst_shape, &transformed_filter)); + if (!FastBoundsCheck(filter.NumElements(), + std::numeric_limits::max())) { + return errors::InvalidArgument( + "Filter tensor num elements (", filter.NumElements(), + ") exceeds 32-bit limit for GPU transformation"); + } functor::TransformFilter()( ctx->eigen_device(), dst_format, To32Bit(filter.tensor()), diff --git a/tensorflow/core/kernels/conv_grad_input_ops_3d.cc b/tensorflow/core/kernels/conv_grad_input_ops_3d.cc index dc991c43c2b726..ba60f11a3807eb 100644 --- a/tensorflow/core/kernels/conv_grad_input_ops_3d.cc +++ b/tensorflow/core/kernels/conv_grad_input_ops_3d.cc @@ -21,6 +21,7 @@ limitations under the License. #include #include +#include "tensorflow/core/framework/bounds_check.h" #include "tensorflow/core/framework/kernel_shape_util.h" #include "tensorflow/core/framework/numeric_op.h" #include "tensorflow/core/framework/op_kernel.h" @@ -877,6 +878,14 @@ void LaunchConvBackpropInputOpImpl( context->allocate_temp(DataTypeToEnum::value, dst_shape, &transformed_filter)); + OP_REQUIRES( + context, + FastBoundsCheck(filter.NumElements(), + std::numeric_limits::max()), + errors::InvalidArgument("Filter tensor num elements (", + filter.NumElements(), + ") exceeds 32-bit limit for GPU transformation")); + functor::TransformFilter()( context->eigen_device(), dst_format, To32Bit(filter.tensor()), diff --git a/tensorflow/core/kernels/conv_ops_fused_impl.h b/tensorflow/core/kernels/conv_ops_fused_impl.h index 2a62c27d1f4a6e..c9a91a45b43b8b 100644 --- a/tensorflow/core/kernels/conv_ops_fused_impl.h +++ b/tensorflow/core/kernels/conv_ops_fused_impl.h @@ -551,6 +551,14 @@ struct LaunchFusedConv2DOp { TF_RETURN_IF_ERROR(context->allocate_temp( DataTypeToEnum::value, dst_shape, &transformed_filter)); + + if (!FastBoundsCheck(filter.NumElements(), + std::numeric_limits::max())) { + return errors::InvalidArgument( + "Filter tensor num elements (", filter.NumElements(), + ") exceeds 32-bit limit for GPU transformation"); + } + functor::TransformFilter()( context->eigen_device(), dst_format, To32Bit(filter.tensor()), diff --git a/tensorflow/core/kernels/conv_ops_impl.h b/tensorflow/core/kernels/conv_ops_impl.h index 5c84dd9bac12ca..66e0102f8c50e8 100644 --- a/tensorflow/core/kernels/conv_ops_impl.h +++ b/tensorflow/core/kernels/conv_ops_impl.h @@ -1137,6 +1137,14 @@ void LaunchConvOpImpl(OpKernelContext* context, bool cudnn_use_autotune, context->allocate_temp(DataTypeToEnum::value, dst_shape, &transformed_filter)); + OP_REQUIRES( + context, + FastBoundsCheck(filter.NumElements(), + std::numeric_limits::max()), + errors::InvalidArgument("Filter tensor num elements (", + filter.NumElements(), + ") exceeds 32-bit limit for GPU transformation")); + // Filter: [(spatial_dims), in, out] (HWIO) // T_filter: [out, in, (spatial_dims)] (OIHW) or // T_filter: [out, (spatial_dims), in] (OHWI) diff --git a/tensorflow/python/kernel_tests/nn_ops/conv_ops_test.py b/tensorflow/python/kernel_tests/nn_ops/conv_ops_test.py index 7d3aa8fdd59c88..295ed9b5fdda07 100644 --- a/tensorflow/python/kernel_tests/nn_ops/conv_ops_test.py +++ b/tensorflow/python/kernel_tests/nn_ops/conv_ops_test.py @@ -3197,6 +3197,24 @@ def testConv2DBackpropInputInvalidOutBackpropRaiseError(self): dilations=[1, 1, 1, 1]) self.evaluate(t) + def testConv2DFilterExceeds32BitLimit(self): + # Verify that filter configurations with invalid element count exceeding + # the 32-bit limit for GPU transformation raise clean errors rather than + # crashing with process abort or illegal memory write. + with self.assertRaises((errors_impl.InvalidArgumentError, ValueError)): + with self.cached_session(): + input_tensor = constant_op.constant( + 1.0, shape=[1, 1, 1, 1], dtype=dtypes.float32) + # 3 * 100 * 8044155 = 2413246500 > INT32_MAX + filter_tensor = constant_op.constant( + 0.0, shape=[3, 100, 8044155, 1], dtype=dtypes.float32) + t = gen_nn_ops.conv2d( + input=input_tensor, + filter=filter_tensor, + strides=[1, 1, 1, 1], + padding="SAME") + self.evaluate(t) + @test_util.run_all_without_tensor_float_32("Avoid TF32 conv on GPU") class DepthwiseConv2DTest(test.TestCase): From b21b554e4863d86eaeac4ffa074bd94e7f26f78d Mon Sep 17 00:00:00 2001 From: Rohith Pariki Date: Thu, 24 Sep 2026 21:48:16 +0530 Subject: [PATCH 13/58] Address PR feedback: Move bounds check before allocation, use direct comparison, and remove memory-heavy python test --- tensorflow/core/kernels/conv_grad_input_ops.cc | 7 +++---- .../core/kernels/conv_grad_input_ops_3d.cc | 13 ++++++------- tensorflow/core/kernels/conv_ops_fused_impl.h | 8 +++----- tensorflow/core/kernels/conv_ops_impl.h | 11 +++++------ .../kernel_tests/nn_ops/conv_ops_test.py | 18 ------------------ 5 files changed, 17 insertions(+), 40 deletions(-) diff --git a/tensorflow/core/kernels/conv_grad_input_ops.cc b/tensorflow/core/kernels/conv_grad_input_ops.cc index 93d5cbae5eba23..4bce81e66cb0ec 100644 --- a/tensorflow/core/kernels/conv_grad_input_ops.cc +++ b/tensorflow/core/kernels/conv_grad_input_ops.cc @@ -291,14 +291,13 @@ void LaunchConv2DBackpropInputOpGpuImpl( : TensorShape({filter.dim_size(3), filter.dim_size(0), filter.dim_size(1), filter.dim_size(2)}); - TF_RETURN_IF_ERROR(ctx->allocate_temp(DataTypeToEnum::value, dst_shape, - &transformed_filter)); - if (!FastBoundsCheck(filter.NumElements(), - std::numeric_limits::max())) { + if (filter.NumElements() > std::numeric_limits::max()) { return errors::InvalidArgument( "Filter tensor num elements (", filter.NumElements(), ") exceeds 32-bit limit for GPU transformation"); } + TF_RETURN_IF_ERROR(ctx->allocate_temp(DataTypeToEnum::value, dst_shape, + &transformed_filter)); functor::TransformFilter()( ctx->eigen_device(), dst_format, To32Bit(filter.tensor()), diff --git a/tensorflow/core/kernels/conv_grad_input_ops_3d.cc b/tensorflow/core/kernels/conv_grad_input_ops_3d.cc index ba60f11a3807eb..222436da1fa02d 100644 --- a/tensorflow/core/kernels/conv_grad_input_ops_3d.cc +++ b/tensorflow/core/kernels/conv_grad_input_ops_3d.cc @@ -21,7 +21,7 @@ limitations under the License. #include #include -#include "tensorflow/core/framework/bounds_check.h" + #include "tensorflow/core/framework/kernel_shape_util.h" #include "tensorflow/core/framework/numeric_op.h" #include "tensorflow/core/framework/op_kernel.h" @@ -874,18 +874,17 @@ void LaunchConvBackpropInputOpImpl( : TensorShape({filter_shape.dim_size(4), dims.filter_size(0), dims.filter_size(1), dims.filter_size(2), filter_shape.dim_size(3)}); - OP_REQUIRES_OK(context, - context->allocate_temp(DataTypeToEnum::value, dst_shape, - &transformed_filter)); - OP_REQUIRES( context, - FastBoundsCheck(filter.NumElements(), - std::numeric_limits::max()), + filter.NumElements() <= std::numeric_limits::max(), errors::InvalidArgument("Filter tensor num elements (", filter.NumElements(), ") exceeds 32-bit limit for GPU transformation")); + OP_REQUIRES_OK(context, + context->allocate_temp(DataTypeToEnum::value, dst_shape, + &transformed_filter)); + functor::TransformFilter()( context->eigen_device(), dst_format, To32Bit(filter.tensor()), diff --git a/tensorflow/core/kernels/conv_ops_fused_impl.h b/tensorflow/core/kernels/conv_ops_fused_impl.h index c9a91a45b43b8b..d010c3eb163027 100644 --- a/tensorflow/core/kernels/conv_ops_fused_impl.h +++ b/tensorflow/core/kernels/conv_ops_fused_impl.h @@ -549,15 +549,13 @@ struct LaunchFusedConv2DOp { : TensorShape({filter.dim_size(3), filter.dim_size(0), filter.dim_size(1), filter.dim_size(2)}); - TF_RETURN_IF_ERROR(context->allocate_temp( - DataTypeToEnum::value, dst_shape, &transformed_filter)); - - if (!FastBoundsCheck(filter.NumElements(), - std::numeric_limits::max())) { + if (filter.NumElements() > std::numeric_limits::max()) { return errors::InvalidArgument( "Filter tensor num elements (", filter.NumElements(), ") exceeds 32-bit limit for GPU transformation"); } + TF_RETURN_IF_ERROR(context->allocate_temp( + DataTypeToEnum::value, dst_shape, &transformed_filter)); functor::TransformFilter()( context->eigen_device(), dst_format, diff --git a/tensorflow/core/kernels/conv_ops_impl.h b/tensorflow/core/kernels/conv_ops_impl.h index 66e0102f8c50e8..26c5c057f7bf2d 100644 --- a/tensorflow/core/kernels/conv_ops_impl.h +++ b/tensorflow/core/kernels/conv_ops_impl.h @@ -1133,18 +1133,17 @@ void LaunchConvOpImpl(OpKernelContext* context, bool cudnn_use_autotune, } } TensorShape dst_shape(dst_shape_vec); - OP_REQUIRES_OK(context, - context->allocate_temp(DataTypeToEnum::value, dst_shape, - &transformed_filter)); - OP_REQUIRES( context, - FastBoundsCheck(filter.NumElements(), - std::numeric_limits::max()), + filter.NumElements() <= std::numeric_limits::max(), errors::InvalidArgument("Filter tensor num elements (", filter.NumElements(), ") exceeds 32-bit limit for GPU transformation")); + OP_REQUIRES_OK(context, + context->allocate_temp(DataTypeToEnum::value, dst_shape, + &transformed_filter)); + // Filter: [(spatial_dims), in, out] (HWIO) // T_filter: [out, in, (spatial_dims)] (OIHW) or // T_filter: [out, (spatial_dims), in] (OHWI) diff --git a/tensorflow/python/kernel_tests/nn_ops/conv_ops_test.py b/tensorflow/python/kernel_tests/nn_ops/conv_ops_test.py index 295ed9b5fdda07..7d3aa8fdd59c88 100644 --- a/tensorflow/python/kernel_tests/nn_ops/conv_ops_test.py +++ b/tensorflow/python/kernel_tests/nn_ops/conv_ops_test.py @@ -3197,24 +3197,6 @@ def testConv2DBackpropInputInvalidOutBackpropRaiseError(self): dilations=[1, 1, 1, 1]) self.evaluate(t) - def testConv2DFilterExceeds32BitLimit(self): - # Verify that filter configurations with invalid element count exceeding - # the 32-bit limit for GPU transformation raise clean errors rather than - # crashing with process abort or illegal memory write. - with self.assertRaises((errors_impl.InvalidArgumentError, ValueError)): - with self.cached_session(): - input_tensor = constant_op.constant( - 1.0, shape=[1, 1, 1, 1], dtype=dtypes.float32) - # 3 * 100 * 8044155 = 2413246500 > INT32_MAX - filter_tensor = constant_op.constant( - 0.0, shape=[3, 100, 8044155, 1], dtype=dtypes.float32) - t = gen_nn_ops.conv2d( - input=input_tensor, - filter=filter_tensor, - strides=[1, 1, 1, 1], - padding="SAME") - self.evaluate(t) - @test_util.run_all_without_tensor_float_32("Avoid TF32 conv on GPU") class DepthwiseConv2DTest(test.TestCase): From 8d78e0875a94c6a025a923efe26a1561c9559450 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Thu, 24 Sep 2026 20:56:24 -0700 Subject: [PATCH 14/58] Hide only the LLVM C API from libtensorflow_framework Hiding the llvm:: and mlir:: C++ namespaces broke the build, because libtensorflow_cc links against those symbols from libtensorflow_framework. It imports 5,351 of them in the 2.21.0 Linux wheel, and 3,816 LLVM symbols in a macOS nightly, where the *llvm* and *LLVM* patterns would have hidden them just the same. The LLVM C API is different: no other TensorFlow library imports any of it, on either platform, so hiding it cannot break a link. Keep only that. On Linux the pattern is LLVM*. On macOS it is _LLVM*, which carries the Mach-O leading underscore and, unlike *LLVM*, does not also match C++ names that mention LLVM types. --- tensorflow/tf_framework_version_script.lds | 22 ++++++---------------- tensorflow/tf_private_symbols.lds | 3 +-- 2 files changed, 7 insertions(+), 18 deletions(-) diff --git a/tensorflow/tf_framework_version_script.lds b/tensorflow/tf_framework_version_script.lds index 9d933db6f770d7..038574de752287 100644 --- a/tensorflow/tf_framework_version_script.lds +++ b/tensorflow/tf_framework_version_script.lds @@ -2,23 +2,13 @@ tensorflow { global: *; - # Statically linked third-party code whose symbols must not reach the - # process-wide resolution scope. LLVM in particular collides with other - # copies loaded into the same process, for example the LLVM that llvmlite - # and numba bring in, which resolve against TensorFlow's build and crash on - # the mismatched layout. + # The C API of the statically linked LLVM. Exported, it lets another LLVM + # loaded into the same process, such as the one llvmlite and numba bring in, + # resolve its C API calls against TensorFlow's copy instead of its own. No + # other TensorFlow library imports these symbols. # - # Matched by namespace rather than by substring so that TensorFlow's own - # symbols are not caught: a plain *mlir* would also hide names such as - # tensorflow::tf_xla_test_use_mlir. TensorFlow's symbols must stay exported, - # because custom op libraries loaded through tf.load_op_library resolve - # REGISTER_OP and REGISTER_KERNEL_BUILDER against this library. - # - # LLVM* covers the LLVM C API, whose symbols are not in the llvm namespace. + # The llvm:: and mlir:: C++ symbols have to stay exported: libtensorflow_cc + # links against them from this library. local: - extern "C++" { - llvm::*; - mlir::*; - }; LLVM*; }; diff --git a/tensorflow/tf_private_symbols.lds b/tensorflow/tf_private_symbols.lds index 488eed51d27869..264a68e0d7f454 100644 --- a/tensorflow/tf_private_symbols.lds +++ b/tensorflow/tf_private_symbols.lds @@ -6,6 +6,5 @@ _jzero_far _jcopy_* _jsimd_* _hwloc_* -*llvm* -*LLVM* +_LLVM* *mlir* From 38dcd10bfa717b60a6423c20e78afbc245071511 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Fri, 25 Sep 2026 08:54:26 -0700 Subject: [PATCH 15/58] Pass check_dtypes=False to assertAllClose in testRot90InvalidAxes --- tensorflow/python/ops/numpy_ops/tests/np_test.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tensorflow/python/ops/numpy_ops/tests/np_test.py b/tensorflow/python/ops/numpy_ops/tests/np_test.py index fd2cc9abdeead0..352e849f1df191 100644 --- a/tensorflow/python/ops/numpy_ops/tests/np_test.py +++ b/tensorflow/python/ops/numpy_ops/tests/np_test.py @@ -2094,7 +2094,7 @@ def testRot90InvalidAxes(self): tnp.rot90(a, axes=(0,)) # In-bounds negative axes remain valid. self.assertAllClose(tnp.rot90(a, axes=(-2, -1)), onp.rot90( - onp.ones((2, 3)), axes=(-2, -1))) + onp.ones((2, 3)), axes=(-2, -1)), check_dtypes=False) # TODO(mattjj): test infix operator overrides From 10a810fc252b8de6052022a3785b7c0ee18c2f09 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Fri, 25 Sep 2026 09:44:54 -0700 Subject: [PATCH 16/58] Keep the LLVM target initializers exported In the CUDA builds, libtensorflow_cc links tfcompile's InitializeTargets(), which calls the LLVMInitialize* functions for each LLVM target and imports them from libtensorflow_framework, so hiding all of LLVM* broke that link. They are the only LLVM C API functions that TensorFlow or XLA code calls. Export them again and keep hiding the rest of the C API. --- tensorflow/tf_framework_version_script.lds | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/tensorflow/tf_framework_version_script.lds b/tensorflow/tf_framework_version_script.lds index 038574de752287..4a4956bb56fefe 100644 --- a/tensorflow/tf_framework_version_script.lds +++ b/tensorflow/tf_framework_version_script.lds @@ -2,10 +2,16 @@ tensorflow { global: *; - # The C API of the statically linked LLVM. Exported, it lets another LLVM - # loaded into the same process, such as the one llvmlite and numba bring in, - # resolve its C API calls against TensorFlow's copy instead of its own. No - # other TensorFlow library imports these symbols. + # The target initializers are the one part of the LLVM C API that + # TensorFlow calls: tfcompile's InitializeTargets(), which is linked into + # libtensorflow_cc, imports them from this library. A global wildcard + # takes precedence over the local LLVM* below. + LLVMInitialize*; + + # The rest of the C API of the statically linked LLVM. Exported, it lets + # another LLVM loaded into the same process, such as the one llvmlite and + # numba bring in, resolve its C API calls against TensorFlow's copy instead + # of its own. # # The llvm:: and mlir:: C++ symbols have to stay exported: libtensorflow_cc # links against them from this library. From 3122d28764cfd1e1b1300fc683e5e2928c94048e Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Sat, 26 Sep 2026 13:43:52 -0700 Subject: [PATCH 17/58] Drop the empty-input early return in ParseAndCheckBoxSizes With the rank checks first, the early return for empty boxes and box_index no longer protects anything: well-formed empty inputs pass the column and row checks on their own. It only let malformed empty boxes through, such as shape [0, 5] or [2, 0], which every crop-and-resize kernel then accepted. Remove it, and extend the test to cover the boxes gradient with a rank-2 box_ind, empty boxes with the wrong number of columns, and well-formed empty inputs to the boxes gradient. --- .../core/kernels/image/crop_and_resize_op.cc | 4 --- tensorflow/python/ops/image_grad_test_base.py | 32 +++++++++++++++++-- 2 files changed, 30 insertions(+), 6 deletions(-) diff --git a/tensorflow/core/kernels/image/crop_and_resize_op.cc b/tensorflow/core/kernels/image/crop_and_resize_op.cc index 978a1ab10b8358..6afb04e4648f36 100644 --- a/tensorflow/core/kernels/image/crop_and_resize_op.cc +++ b/tensorflow/core/kernels/image/crop_and_resize_op.cc @@ -66,10 +66,6 @@ static inline absl::Status ParseAndCheckBoxSizes(const Tensor& boxes, return absl::InvalidArgumentError(absl::StrCat( "box_index must be 1-D, got ", box_index.shape().DebugString())); } - if (boxes.NumElements() == 0 && box_index.NumElements() == 0) { - *num_boxes = 0; - return absl::OkStatus(); - } *num_boxes = boxes.dim_size(0); if (boxes.dim_size(1) != 4) { return absl::InvalidArgumentError("boxes must have 4 columns"); diff --git a/tensorflow/python/ops/image_grad_test_base.py b/tensorflow/python/ops/image_grad_test_base.py index 65f5c7f157fb0e..aaffff464568a3 100644 --- a/tensorflow/python/ops/image_grad_test_base.py +++ b/tensorflow/python/ops/image_grad_test_base.py @@ -453,6 +453,7 @@ def testMalformedEmptyBoxesRaisesError(self): # rank-2 empty `box_ind` used to pass validation and crash the process # with a fatal CHECK failure instead of raising InvalidArgumentError. grads = array_ops.zeros([0, 1, 1, 1], dtype=dtypes.float32) + image = array_ops.zeros([2, 7, 7, 1], dtype=dtypes.float32) image_size = constant_op.constant([2, 7, 7, 1], dtype=dtypes.int32) valid_boxes = array_ops.zeros([0, 4], dtype=dtypes.float32) valid_box_ind = array_ops.zeros([0], dtype=dtypes.int32) @@ -475,15 +476,35 @@ def testMalformedEmptyBoxesRaisesError(self): box_ind=array_ops.zeros([0, 0], dtype=dtypes.int32), image_size=image_size, T=dtypes.float32)) + # Empty boxes with the wrong number of columns must be rejected too. + with self.assertRaisesRegex( + (errors_impl.InvalidArgumentError, ValueError), "4 columns|must be 4"): + self.evaluate( + gen_image_ops.crop_and_resize_grad_image( + grads=grads, + boxes=array_ops.zeros([0, 5], dtype=dtypes.float32), + box_ind=valid_box_ind, + image_size=image_size, + T=dtypes.float32)) with self.assertRaisesRegex( (errors_impl.InvalidArgumentError, ValueError), "boxes must be 2-D"): self.evaluate( gen_image_ops.crop_and_resize_grad_boxes( grads=grads, - image=array_ops.zeros([2, 7, 7, 1], dtype=dtypes.float32), + image=image, boxes=array_ops.zeros([0], dtype=dtypes.float32), box_ind=valid_box_ind)) - # Well-formed empty boxes must keep working. + # Rank-2 empty box_ind for the boxes gradient. + with self.assertRaisesRegex( + (errors_impl.InvalidArgumentError, ValueError), "box_index must be 1-D" + ): + self.evaluate( + gen_image_ops.crop_and_resize_grad_boxes( + grads=grads, + image=image, + boxes=valid_boxes, + box_ind=array_ops.zeros([0, 0], dtype=dtypes.int32))) + # Well-formed empty inputs must keep working. output = self.evaluate( gen_image_ops.crop_and_resize_grad_image( grads=grads, @@ -492,6 +513,13 @@ def testMalformedEmptyBoxesRaisesError(self): image_size=image_size, T=dtypes.float32)) self.assertEqual((2, 7, 7, 1), output.shape) + output = self.evaluate( + gen_image_ops.crop_and_resize_grad_boxes( + grads=grads, + image=image, + boxes=valid_boxes, + box_ind=valid_box_ind)) + self.assertEqual((0, 4), output.shape) def _randomUniformAvoidAnchors(self, low, high, anchors, radius, num_samples): """Generate samples that are far enough from a set of anchor points. From 668a864425a139c550f7fe7052b3b2ac7502d709 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Sat, 26 Sep 2026 15:00:55 -0700 Subject: [PATCH 18/58] Format the crop-and-resize changes Run clang-format on ParseAndCheckBoxSizes and pyink on testMalformedEmptyBoxesRaisesError, the code this change touches. No functional change. --- .../core/kernels/image/crop_and_resize_op.cc | 4 +- tensorflow/python/ops/image_grad_test_base.py | 44 ++++++++++++------- 2 files changed, 31 insertions(+), 17 deletions(-) diff --git a/tensorflow/core/kernels/image/crop_and_resize_op.cc b/tensorflow/core/kernels/image/crop_and_resize_op.cc index 6afb04e4648f36..4f9063471dd510 100644 --- a/tensorflow/core/kernels/image/crop_and_resize_op.cc +++ b/tensorflow/core/kernels/image/crop_and_resize_op.cc @@ -59,8 +59,8 @@ static inline absl::Status ParseAndCheckBoxSizes(const Tensor& boxes, // [num_boxes]. The ranks must be validated even when both tensors are // empty, since the kernels later access them as rank-2 and rank-1 tensors. if (boxes.dims() != 2) { - return absl::InvalidArgumentError(absl::StrCat( - "boxes must be 2-D, got ", boxes.shape().DebugString())); + return absl::InvalidArgumentError( + absl::StrCat("boxes must be 2-D, got ", boxes.shape().DebugString())); } if (box_index.dims() != 1) { return absl::InvalidArgumentError(absl::StrCat( diff --git a/tensorflow/python/ops/image_grad_test_base.py b/tensorflow/python/ops/image_grad_test_base.py index aaffff464568a3..ccb2cbcc8c9358 100644 --- a/tensorflow/python/ops/image_grad_test_base.py +++ b/tensorflow/python/ops/image_grad_test_base.py @@ -458,16 +458,19 @@ def testMalformedEmptyBoxesRaisesError(self): valid_boxes = array_ops.zeros([0, 4], dtype=dtypes.float32) valid_box_ind = array_ops.zeros([0], dtype=dtypes.int32) with self.assertRaisesRegex( - (errors_impl.InvalidArgumentError, ValueError), "boxes must be 2-D"): + (errors_impl.InvalidArgumentError, ValueError), 'boxes must be 2-D' + ): self.evaluate( gen_image_ops.crop_and_resize_grad_image( grads=grads, boxes=array_ops.zeros([0], dtype=dtypes.float32), box_ind=valid_box_ind, image_size=image_size, - T=dtypes.float32)) + T=dtypes.float32, + ) + ) with self.assertRaisesRegex( - (errors_impl.InvalidArgumentError, ValueError), "box_index must be 1-D" + (errors_impl.InvalidArgumentError, ValueError), 'box_index must be 1-D' ): self.evaluate( gen_image_ops.crop_and_resize_grad_image( @@ -475,35 +478,45 @@ def testMalformedEmptyBoxesRaisesError(self): boxes=valid_boxes, box_ind=array_ops.zeros([0, 0], dtype=dtypes.int32), image_size=image_size, - T=dtypes.float32)) + T=dtypes.float32, + ) + ) # Empty boxes with the wrong number of columns must be rejected too. with self.assertRaisesRegex( - (errors_impl.InvalidArgumentError, ValueError), "4 columns|must be 4"): + (errors_impl.InvalidArgumentError, ValueError), '4 columns|must be 4' + ): self.evaluate( gen_image_ops.crop_and_resize_grad_image( grads=grads, boxes=array_ops.zeros([0, 5], dtype=dtypes.float32), box_ind=valid_box_ind, image_size=image_size, - T=dtypes.float32)) + T=dtypes.float32, + ) + ) with self.assertRaisesRegex( - (errors_impl.InvalidArgumentError, ValueError), "boxes must be 2-D"): + (errors_impl.InvalidArgumentError, ValueError), 'boxes must be 2-D' + ): self.evaluate( gen_image_ops.crop_and_resize_grad_boxes( grads=grads, image=image, boxes=array_ops.zeros([0], dtype=dtypes.float32), - box_ind=valid_box_ind)) + box_ind=valid_box_ind, + ) + ) # Rank-2 empty box_ind for the boxes gradient. with self.assertRaisesRegex( - (errors_impl.InvalidArgumentError, ValueError), "box_index must be 1-D" + (errors_impl.InvalidArgumentError, ValueError), 'box_index must be 1-D' ): self.evaluate( gen_image_ops.crop_and_resize_grad_boxes( grads=grads, image=image, boxes=valid_boxes, - box_ind=array_ops.zeros([0, 0], dtype=dtypes.int32))) + box_ind=array_ops.zeros([0, 0], dtype=dtypes.int32), + ) + ) # Well-formed empty inputs must keep working. output = self.evaluate( gen_image_ops.crop_and_resize_grad_image( @@ -511,14 +524,15 @@ def testMalformedEmptyBoxesRaisesError(self): boxes=valid_boxes, box_ind=valid_box_ind, image_size=image_size, - T=dtypes.float32)) + T=dtypes.float32, + ) + ) self.assertEqual((2, 7, 7, 1), output.shape) output = self.evaluate( gen_image_ops.crop_and_resize_grad_boxes( - grads=grads, - image=image, - boxes=valid_boxes, - box_ind=valid_box_ind)) + grads=grads, image=image, boxes=valid_boxes, box_ind=valid_box_ind + ) + ) self.assertEqual((0, 4), output.shape) def _randomUniformAvoidAnchors(self, low, high, anchors, radius, num_samples): From 71136236dd60d67883a4505f24ba7d278e72d422 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Sat, 26 Sep 2026 21:56:13 -0700 Subject: [PATCH 19/58] Validate axes in experimental.numpy transpose --- .../python/ops/numpy_ops/np_array_ops.py | 26 +++++++++++++++++++ .../python/ops/numpy_ops/np_array_ops_test.py | 16 ++++++++++++ 2 files changed, 42 insertions(+) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index 9811798a45bd13..7c7ca50355bbf0 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -963,6 +963,32 @@ def flatten(a, order='C'): @np_utils.np_doc('transpose') def transpose(a, axes=None): a = asarray(a) + + maybe_rank = a.shape.rank + if maybe_rank is not None and isinstance(axes, (tuple, list)): + # Match np.transpose behavior: raise a ValueError at trace time for + # invalid `axes` instead of letting the underlying op produce an + # opaque error deeper in the stack. + if len(axes) != maybe_rank: + raise ValueError( + f'axes don\'t match array. Expected {maybe_rank} axes, got ' + f'{len(axes)}.' + ) + normalized_axes = [] + for ax in axes: + if isinstance(ax, (int, np.integer)): + normalized = ax + maybe_rank if ax < 0 else ax + if normalized < 0 or normalized >= maybe_rank: + raise ValueError( + f'Argument `axes` (received axes={ax}) is out of bounds for ' + f'input {a} of rank {maybe_rank}.' + ) + normalized_axes.append(normalized) + if len(normalized_axes) == maybe_rank and len(set(normalized_axes)) != len( + normalized_axes + ): + raise ValueError('repeated axis in transpose') + if axes is not None: axes = asarray(axes) return array_ops.transpose(a=a, perm=axes) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops_test.py b/tensorflow/python/ops/numpy_ops/np_array_ops_test.py index 89e0fc6afb3067..658935bac283d9 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops_test.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops_test.py @@ -1202,6 +1202,22 @@ def run_test(arr, axes=None): run_test(np.arange(30).reshape(2, 3, 5).tolist(), [1, 2, 0]) run_test(np.arange(30).reshape(2, 3, 5).tolist(), [2, 0, 1]) run_test(np.arange(30).reshape(2, 3, 5).tolist(), [2, 1, 0]) + a = np.arange(6).reshape(2, 3) + # Valid negative and mixed axes still work. + self.match( + np_array_ops.transpose(a, [0, -1]), np.transpose(a, [0, -1])) + self.match( + np_array_ops.transpose(a, [-2, -1]), np.transpose(a, [-2, -1])) + with self.assertRaisesRegex(ValueError, "axes don't match array"): + np_array_ops.transpose(a, [0, 1, 2]) + with self.assertRaisesRegex(ValueError, 'out of bounds'): + np_array_ops.transpose(a, [0, 2]) + with self.assertRaisesRegex(ValueError, 'out of bounds'): + np_array_ops.transpose(a, [0, -3]) + with self.assertRaisesRegex(ValueError, 'repeated axis'): + np_array_ops.transpose(a, [1, 1]) + with self.assertRaisesRegex(ValueError, 'repeated axis'): + np_array_ops.transpose(a, [0, -2]) def match_shape(self, actual, expected, msg=None): if msg: From 2d8688d75d14490e923452681276f478b84fea45 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Mon, 28 Sep 2026 08:39:15 -0700 Subject: [PATCH 20/58] Address review: allocation-free duplicate check, fix mixed-type bypass, add rank-0/1 tests, quote style --- .../python/ops/numpy_ops/np_array_ops.py | 27 ++++++++++++------- .../python/ops/numpy_ops/np_array_ops_test.py | 14 ++++++++++ 2 files changed, 31 insertions(+), 10 deletions(-) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index 7c7ca50355bbf0..b5e0a415758f3c 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -968,25 +968,32 @@ def transpose(a, axes=None): if maybe_rank is not None and isinstance(axes, (tuple, list)): # Match np.transpose behavior: raise a ValueError at trace time for # invalid `axes` instead of letting the underlying op produce an - # opaque error deeper in the stack. + # opaque error deeper in the stack. Duplicate detection uses a single + # integer bitmask (no intermediate list/set allocations), and only + # runs when the static rank and Python-int entries are known. if len(axes) != maybe_rank: raise ValueError( - f'axes don\'t match array. Expected {maybe_rank} axes, got ' - f'{len(axes)}.' + f"axes don't match array. Expected {maybe_rank} axes, got " + f"{len(axes)}." ) - normalized_axes = [] + normalized_mask = 0 + seen_int_axes = 0 for ax in axes: if isinstance(ax, (int, np.integer)): normalized = ax + maybe_rank if ax < 0 else ax if normalized < 0 or normalized >= maybe_rank: raise ValueError( - f'Argument `axes` (received axes={ax}) is out of bounds for ' - f'input {a} of rank {maybe_rank}.' + f"Argument 'axes' (received axes={ax}) is out of bounds for " + f"array of rank {maybe_rank}." ) - normalized_axes.append(normalized) - if len(normalized_axes) == maybe_rank and len(set(normalized_axes)) != len( - normalized_axes - ): + bit = 1 << normalized + if normalized_mask & bit: + raise ValueError('repeated axis in transpose') + normalized_mask |= bit + seen_int_axes += 1 + if seen_int_axes == maybe_rank and normalized_mask != ( + 1 << maybe_rank + ) - 1: raise ValueError('repeated axis in transpose') if axes is not None: diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops_test.py b/tensorflow/python/ops/numpy_ops/np_array_ops_test.py index 658935bac283d9..e3c6669595bb6d 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops_test.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops_test.py @@ -1218,6 +1218,20 @@ def run_test(arr, axes=None): np_array_ops.transpose(a, [1, 1]) with self.assertRaisesRegex(ValueError, 'repeated axis'): np_array_ops.transpose(a, [0, -2]) + # Test scalar rank. + a_scalar = np.array(5) + self.match( + np_array_ops.transpose(a_scalar, []), np.transpose(a_scalar, [])) + with self.assertRaisesRegex(ValueError, "axes don't match array"): + np_array_ops.transpose(a_scalar, [0]) + # Test vector rank. + a_vector = np.array([1, 2, 3]) + with self.assertRaisesRegex(ValueError, "axes don't match array"): + np_array_ops.transpose(a_vector, [0, 1]) + # Duplicate detection must not be bypassed by mixed int/Tensor axes. + a3 = np.arange(24).reshape(2, 3, 4) + with self.assertRaisesRegex(ValueError, 'repeated axis'): + np_array_ops.transpose(a3, [0, 0, constant_op.constant(2)]) def match_shape(self, actual, expected, msg=None): if msg: From 05007937b01a7fee4053aadc151b9ca0bece29aa Mon Sep 17 00:00:00 2001 From: Rohith Pariki Date: Tue, 29 Sep 2026 01:04:24 +0530 Subject: [PATCH 21/58] Address silent early return review: revert TF_PREDICT_FALSE guards, add OP_REQUIRES on callers Per maintainer review (dmiltr3): - Revert if (TF_PREDICT_FALSE(out.size() <= 0)) return; from TransformFilter and ReverseTransformFilter in conv_2d_gpu.h. These caused silent skips with uninitialized output buffers being passed to cuDNN. - Add upfront OP_REQUIRES bounds checks in conv_grad_filter_ops_launcher.cc and conv_grad_filter_ops_3d.cc before allocate_temp, so invalid filters produce a loud InvalidArgument error rather than silent data corruption. Fixes: https://github.com/tensorflow/tensorflow/issues/115734 --- tensorflow/core/kernels/conv_2d_gpu.h | 7 ------- tensorflow/core/kernels/conv_grad_filter_ops_3d.cc | 9 +++++++++ tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc | 9 +++++++++ 3 files changed, 18 insertions(+), 7 deletions(-) diff --git a/tensorflow/core/kernels/conv_2d_gpu.h b/tensorflow/core/kernels/conv_2d_gpu.h index e33294e18546cb..3c9ccb5e6330f8 100644 --- a/tensorflow/core/kernels/conv_2d_gpu.h +++ b/tensorflow/core/kernels/conv_2d_gpu.h @@ -491,9 +491,6 @@ struct TransformFilter { } combined_dims[1] = in.dimension(NDIMS - 2); // input filters combined_dims[2] = in.dimension(NDIMS - 1); // output filters - if (TF_PREDICT_FALSE(out.size() <= 0)) { - return; - } GpuLaunchConfig config = GetGpuLaunchConfig(out.size(), d); if (dst_filter_format == FORMAT_OIHW) { @@ -525,10 +522,6 @@ struct ReverseTransformFilter { typename TTypes::Tensor out) { Dimension<3> combined_dims; - if (TF_PREDICT_FALSE(out.size() <= 0)) { - return; - } - if (src_filter_format == FORMAT_OIHW) { combined_dims[0] = in.dimension(0); // output filters combined_dims[1] = in.dimension(1); // input filters diff --git a/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc b/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc index f657da6337b410..e6168b6630f60a 100644 --- a/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc +++ b/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc @@ -880,6 +880,15 @@ void LaunchConvBackpropFilterOpImpl( : TensorShape({filter_shape.dim_size(4), dims.filter_size(0), dims.filter_size(1), dims.filter_size(2), filter_shape.dim_size(3)}); + // Validate filter element count before allocation to prevent OOM on invalid + // inputs. GPU transformation uses 32-bit indexing via To32Bit(). + OP_REQUIRES(context, + filter_backprop->NumElements() <= + std::numeric_limits::max(), + errors::InvalidArgument( + "Filter tensor num elements (", + filter_backprop->NumElements(), + ") exceeds 32-bit limit for GPU transformation")); OP_REQUIRES_OK(context, context->allocate_temp(DataTypeToEnum::value, dst_shape, &pre_transformed_filter_backprop)); diff --git a/tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc b/tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc index a79772051e94af..3811cae0d2bf58 100644 --- a/tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc +++ b/tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc @@ -399,6 +399,15 @@ void LaunchConv2DBackpropFilterOpImpl( // We compute filter backprop into temporary tensor, and then convert it to // the HWIO data format at the end. + // Validate filter element count before allocation to prevent OOM on invalid + // inputs. GPU transformation uses 32-bit indexing via To32Bit(). + OP_REQUIRES(ctx, + filter_backprop->NumElements() <= + std::numeric_limits::max(), + errors::InvalidArgument( + "Filter tensor num elements (", filter_backprop->NumElements(), + ") exceeds 32-bit limit for GPU transformation")); + Tensor pre_transformed_filter_backprop; OP_REQUIRES_OK( ctx, From d2eec6ad12f434bbb2c5370971230ec563db957d Mon Sep 17 00:00:00 2001 From: Rohith Pariki Date: Tue, 29 Sep 2026 01:54:31 +0530 Subject: [PATCH 22/58] Explicitly include for IWYU and fix formatting Per maintainer review (dmiltr3): - Add explicit #include in conv_grad_filter_ops_3d.cc, conv_grad_filter_ops_launcher.cc, conv_grad_input_ops.cc, conv_grad_input_ops_3d.cc, and conv_ops_fused_impl.h for IWYU compliance. - Remove stray newline before kernel_shape_util.h include in conv_grad_input_ops_3d.cc. --- tensorflow/core/kernels/conv_grad_filter_ops_3d.cc | 1 + tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc | 1 + tensorflow/core/kernels/conv_grad_input_ops.cc | 1 + tensorflow/core/kernels/conv_grad_input_ops_3d.cc | 2 +- tensorflow/core/kernels/conv_ops_fused_impl.h | 1 + 5 files changed, 5 insertions(+), 1 deletion(-) diff --git a/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc b/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc index e6168b6630f60a..c8349670567f45 100644 --- a/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc +++ b/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc @@ -17,6 +17,7 @@ limitations under the License. #define EIGEN_USE_THREADS #include +#include #include #include #include diff --git a/tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc b/tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc index 3811cae0d2bf58..e58591f154dd55 100644 --- a/tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc +++ b/tensorflow/core/kernels/conv_grad_filter_ops_launcher.cc @@ -17,6 +17,7 @@ limitations under the License. #define EIGEN_USE_THREADS #include +#include #include #include diff --git a/tensorflow/core/kernels/conv_grad_input_ops.cc b/tensorflow/core/kernels/conv_grad_input_ops.cc index 4bce81e66cb0ec..93dd9b7a7063cf 100644 --- a/tensorflow/core/kernels/conv_grad_input_ops.cc +++ b/tensorflow/core/kernels/conv_grad_input_ops.cc @@ -17,6 +17,7 @@ limitations under the License. #include "tensorflow/core/kernels/conv_grad_input_ops.h" +#include #include #include "tensorflow/core/profiler/lib/scoped_annotation.h" diff --git a/tensorflow/core/kernels/conv_grad_input_ops_3d.cc b/tensorflow/core/kernels/conv_grad_input_ops_3d.cc index 222436da1fa02d..bfcdb9fc38e868 100644 --- a/tensorflow/core/kernels/conv_grad_input_ops_3d.cc +++ b/tensorflow/core/kernels/conv_grad_input_ops_3d.cc @@ -17,11 +17,11 @@ limitations under the License. #define EIGEN_USE_THREADS #include +#include #include #include #include - #include "tensorflow/core/framework/kernel_shape_util.h" #include "tensorflow/core/framework/numeric_op.h" #include "tensorflow/core/framework/op_kernel.h" diff --git a/tensorflow/core/kernels/conv_ops_fused_impl.h b/tensorflow/core/kernels/conv_ops_fused_impl.h index d010c3eb163027..1b47bbe60198d9 100644 --- a/tensorflow/core/kernels/conv_ops_fused_impl.h +++ b/tensorflow/core/kernels/conv_ops_fused_impl.h @@ -38,6 +38,7 @@ limitations under the License. #define EIGEN_USE_GPU #endif // GOOGLE_CUDA +#include #include #include #include From c4bec43a89600c9afd0c5d01e580f7a2cc4cb640 Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Mon, 28 Sep 2026 13:51:12 -0700 Subject: [PATCH 23/58] Reject negative split sizes before summing them in SplitV The overflow guard kept determined_size itself in range, but a large negative size next to a -1, such as [-1, INT64_MIN], still made input_size_split_dim - determined_size overflow when computing the -1 size. Reject a negative size before summing it, as the later check already did for the error message, so that 0 <= determined_size <= input_size_split_dim and only the upper bound needs a guard. Extend testSizeSplitsOverflowRaises to int32 split sizes and to an overflow next to a -1, accept ValueError for graph mode, and call array_ops.split instead of the generated op, which drops the array_ops_gen dependency this change had added. --- tensorflow/core/kernels/split_v_op.cc | 15 ++++---- .../python/kernel_tests/array_ops/BUILD | 1 - .../kernel_tests/array_ops/split_op_test.py | 34 ++++++++++++------- 3 files changed, 31 insertions(+), 19 deletions(-) diff --git a/tensorflow/core/kernels/split_v_op.cc b/tensorflow/core/kernels/split_v_op.cc index 640963f981bfd2..5e53927a508df2 100644 --- a/tensorflow/core/kernels/split_v_op.cc +++ b/tensorflow/core/kernels/split_v_op.cc @@ -129,18 +129,21 @@ class SplitVOpBase : public OpKernel { "input.")); neg_one_dim = d; } else { + // Reject a negative size before summing it, so that + // 0 <= determined_size <= input_size_split_dim below. Otherwise a + // large negative size, such as the minimum of Tlen next to a -1, makes + // `input_size_split_dim - determined_size` overflow. + OP_REQUIRES(context, size >= 0, + errors::InvalidArgument("Split size at index ", d, + " must be >= 0. Got: ", size)); // Accumulate with an explicit overflow guard. `determined_size += size` // wraps for large `size_splits`, and a wrapped total can equal // `input_size_split_dim` and pass the check below, letting the aligned // slicing path compute an endpoint that reaches a fatal `Tensor::Slice` // invariant. Rejecting overflow here also keeps that later path safe, // since the accepted total then bounds every partial sum. - const bool overflow = - (size > 0 && - determined_size > std::numeric_limits::max() - size) || - (size < 0 && - determined_size < std::numeric_limits::min() - size); - OP_REQUIRES(context, !overflow, + OP_REQUIRES(context, + determined_size <= std::numeric_limits::max() - size, absl::InvalidArgumentError(absl::StrCat( "Sum of size_splits overflows the index type at index ", d, "."))); diff --git a/tensorflow/python/kernel_tests/array_ops/BUILD b/tensorflow/python/kernel_tests/array_ops/BUILD index 40f9164b556103..1278cc3cb75be4 100644 --- a/tensorflow/python/kernel_tests/array_ops/BUILD +++ b/tensorflow/python/kernel_tests/array_ops/BUILD @@ -779,7 +779,6 @@ cuda_py_strict_test( "//tensorflow/python/framework:for_generated_wrappers", "//tensorflow/python/framework:test_lib", "//tensorflow/python/ops:array_ops", - "//tensorflow/python/ops:array_ops_gen", "//tensorflow/python/ops:gradients_impl", "//tensorflow/python/ops:math_ops", "//tensorflow/python/platform:client_testlib", diff --git a/tensorflow/python/kernel_tests/array_ops/split_op_test.py b/tensorflow/python/kernel_tests/array_ops/split_op_test.py index 0449952dd11f61..fcce45ac5726e4 100644 --- a/tensorflow/python/kernel_tests/array_ops/split_op_test.py +++ b/tensorflow/python/kernel_tests/array_ops/split_op_test.py @@ -22,7 +22,6 @@ from tensorflow.python.framework import ops from tensorflow.python.framework import test_util from tensorflow.python.ops import array_ops -from tensorflow.python.ops import gen_array_ops from tensorflow.python.ops import gradients_impl from tensorflow.python.ops import math_ops from tensorflow.python.platform import test @@ -132,17 +131,28 @@ def testSizeSplitsOverflowRaises(self): # fatal `Tensor::Slice` invariant in the aligned slicing path. It must # raise instead. i64_max = (1 << 63) - 1 - value = array_ops.reshape( - constant_op.constant([], dtype=dtypes.float32), - constant_op.constant([i64_max, 0], dtype=dtypes.int64)) - with self.assertRaisesRegex(errors_impl.InvalidArgumentError, "overflow"): - self.evaluate( - gen_array_ops.SplitV( - value=value, - size_splits=constant_op.constant( - [i64_max, i64_max, i64_max, 2], dtype=dtypes.int64), - axis=constant_op.constant(0, dtype=dtypes.int32), - num_split=4)) + i32_max = (1 << 31) - 1 + for input_size, size_splits, dtype in ( + (i64_max, [i64_max, i64_max, i64_max, 2], dtypes.int64), + # A -1 does not keep the other sizes from overflowing. + (i64_max, [-1, i64_max, i64_max, 2], dtypes.int64), + # int32 sizes overflow at their own width. The shape function sums in + # int64, so an input of that size lets graph mode reach the kernel. + (2 * i32_max + 5, [i32_max, i32_max, 5], dtypes.int32), + ): + with self.subTest(size_splits=size_splits, dtype=dtype.name): + value = array_ops.reshape( + constant_op.constant([], dtype=dtypes.float32), + constant_op.constant([input_size, 0], dtype=dtypes.int64), + ) + with self.assertRaisesRegex( + (ValueError, errors_impl.InvalidArgumentError), "overflow" + ): + self.evaluate( + array_ops.split( + value, constant_op.constant(size_splits, dtype=dtype), axis=0 + ) + ) @test_util.run_in_graph_and_eager_modes def testListOfScalarTensors(self): From 52cf0dfb8df9b16915e335cc97260a47a0e0472c Mon Sep 17 00:00:00 2001 From: Vishwak Thatikonda Date: Mon, 28 Sep 2026 15:37:40 -0700 Subject: [PATCH 24/58] Check the SplitV input size against Tlen, and drop the redundant loop input_shape.dim_size(split_dim) was converted to Tlen implicitly, so with int32 or int8 size_splits an input size above the maximum of Tlen was truncated: int32 sizes [1, 2] matched an input of size 2**32 + 3, and the split silently dropped the rest of the input. Check the size in int64 before converting it. Split sizes are now checked for being non-negative before they are summed, and the -1 size is bounded by the input size, so the later loop that checked every size again can no longer fail; remove it. Also build the overflow error with errors::InvalidArgument, as the rest of the file does. The int32 case of testSizeSplitsOverflowRaises used an input larger than INT32_MAX to reach the kernel in graph mode, which the new check rejects first. Use a small input, where graph mode stops at the shape function instead, and test the truncation separately. --- tensorflow/core/kernels/split_v_op.cc | 21 +++++----- .../kernel_tests/array_ops/split_op_test.py | 39 +++++++++++++++---- 2 files changed, 42 insertions(+), 18 deletions(-) diff --git a/tensorflow/core/kernels/split_v_op.cc b/tensorflow/core/kernels/split_v_op.cc index 5e53927a508df2..d14e685c5fd679 100644 --- a/tensorflow/core/kernels/split_v_op.cc +++ b/tensorflow/core/kernels/split_v_op.cc @@ -101,7 +101,15 @@ class SplitVOpBase : public OpKernel { "-input rank(-", input.dims(), ") <= split_dim < input rank (", input.dims(), "), but got ", split_dim_orig))); - Tlen input_size_split_dim = input_shape.dim_size(split_dim); + // Check that the input size fits in Tlen before converting it. Otherwise + // an int32 or int8 Tlen truncates it, and split sizes that sum to the + // truncated size silently drop the rest of the input. + const int64_t actual_input_size = input_shape.dim_size(split_dim); + OP_REQUIRES(context, actual_input_size <= std::numeric_limits::max(), + errors::InvalidArgument( + "Input size along split_dim must be <= max(Tlen). Got: ", + actual_input_size)); + Tlen input_size_split_dim = static_cast(actual_input_size); // Special case 1: num_split == 1. Nothing to do. if (num_split == 1) { @@ -144,9 +152,9 @@ class SplitVOpBase : public OpKernel { // since the accepted total then bounds every partial sum. OP_REQUIRES(context, determined_size <= std::numeric_limits::max() - size, - absl::InvalidArgumentError(absl::StrCat( + errors::InvalidArgument( "Sum of size_splits overflows the index type at index ", - d, "."))); + d, ".")); determined_size += size; } } @@ -166,13 +174,6 @@ class SplitVOpBase : public OpKernel { (*split_sizes_vec)[neg_one_dim] = input_size_split_dim - determined_size; } - for (int i = 0; i < split_sizes_vec->size(); ++i) { - const Tlen& split_size = (*split_sizes_vec)[i]; - OP_REQUIRES(context, split_size >= Tlen(0), - errors::InvalidArgument("Split size at index ", i, - " must be >= 0. Got: ", split_size)); - } - // Special case 2: split along the 1st dimension. The requirements are that // either we are splitting the outer dimension of two or more such that // every outer subpart is aligned or that the split sizes mean that they are diff --git a/tensorflow/python/kernel_tests/array_ops/split_op_test.py b/tensorflow/python/kernel_tests/array_ops/split_op_test.py index fcce45ac5726e4..1a099db8dd502c 100644 --- a/tensorflow/python/kernel_tests/array_ops/split_op_test.py +++ b/tensorflow/python/kernel_tests/array_ops/split_op_test.py @@ -123,7 +123,8 @@ def testExplicitNum(self): @test_util.run_in_graph_and_eager_modes @test_util.disable_xla( "XLA shape inference rejects the reshape to an INT64_MAX dimension, so " - "the test cannot reach the SplitV kernel under XLA") + "the test cannot reach the SplitV kernel under XLA" + ) def testSizeSplitsOverflowRaises(self): # Regression test for GitHub issue 126126. The cumulative sum of # size_splits was computed with unchecked signed addition, so a wrapped @@ -132,13 +133,14 @@ def testSizeSplitsOverflowRaises(self): # raise instead. i64_max = (1 << 63) - 1 i32_max = (1 << 31) - 1 - for input_size, size_splits, dtype in ( - (i64_max, [i64_max, i64_max, i64_max, 2], dtypes.int64), + for input_size, size_splits, dtype, message in ( + (i64_max, [i64_max, i64_max, i64_max, 2], dtypes.int64, "overflow"), # A -1 does not keep the other sizes from overflowing. - (i64_max, [-1, i64_max, i64_max, 2], dtypes.int64), - # int32 sizes overflow at their own width. The shape function sums in - # int64, so an input of that size lets graph mode reach the kernel. - (2 * i32_max + 5, [i32_max, i32_max, 5], dtypes.int32), + (i64_max, [-1, i64_max, i64_max, 2], dtypes.int64, "overflow"), + # int32 sizes overflow at their own width in the kernel. In graph + # mode the shape function, which sums in int64, rejects the mismatch + # before the kernel runs. + (5, [i32_max, i32_max, 5], dtypes.int32, "overflow|can't split axis"), ): with self.subTest(size_splits=size_splits, dtype=dtype.name): value = array_ops.reshape( @@ -146,7 +148,7 @@ def testSizeSplitsOverflowRaises(self): constant_op.constant([input_size, 0], dtype=dtypes.int64), ) with self.assertRaisesRegex( - (ValueError, errors_impl.InvalidArgumentError), "overflow" + (ValueError, errors_impl.InvalidArgumentError), message ): self.evaluate( array_ops.split( @@ -154,6 +156,27 @@ def testSizeSplitsOverflowRaises(self): ) ) + @test_util.run_in_graph_and_eager_modes + @test_util.disable_xla("Checks the SplitV kernel, which XLA replaces") + def testInputSizeAboveSizeSplitsTypeRaises(self): + # With int32 size_splits, an input size above INT32_MAX was truncated to + # int32, so [1, 2] matched an input of size 2**32 + 3 and the split + # silently dropped the rest of the input. In graph mode the shape + # function, which compares in int64, rejects the mismatch first. + value = array_ops.reshape( + constant_op.constant([], dtype=dtypes.float32), + constant_op.constant([(1 << 32) + 3, 0], dtype=dtypes.int64), + ) + with self.assertRaisesRegex( + (ValueError, errors_impl.InvalidArgumentError), + r"must be <= max\(Tlen\)|can't split axis", + ): + self.evaluate( + array_ops.split( + value, constant_op.constant([1, 2], dtype=dtypes.int32), axis=0 + ) + ) + @test_util.run_in_graph_and_eager_modes def testListOfScalarTensors(self): a = math_ops.cast(5, dtypes.int32) From c28d047ecf209e4e15f2c59c46a44fb40d353a61 Mon Sep 17 00:00:00 2001 From: Maddipatla Chatan Date: Tue, 29 Sep 2026 06:54:52 +0530 Subject: [PATCH 25/58] Correct log fragment extraction and main function call Fix log fragment extraction and ensure main function is called correctly. --- ci/official/utilities/extract_resultstore_links.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ci/official/utilities/extract_resultstore_links.py b/ci/official/utilities/extract_resultstore_links.py index 2bd96e1811c171..f0e2a59bd6d13c 100644 --- a/ci/official/utilities/extract_resultstore_links.py +++ b/ci/official/utilities/extract_resultstore_links.py @@ -115,7 +115,7 @@ def parse_log(file_path: str, tests_failed = re.search(TESTS_FAILED_RE, backtrack_line) if build_failed or tests_failed: log_fragment = '\n'.join( - log_lines[max(k - 20, 0):min(end_line + 1, len(log_lines) - 1)]) + log_lines[max(k - 20, 0):min(end_line + 1, len(log_lines))]) lines['log_fragment'] = log_fragment lines['status'] = (InvokeStatus.build_failed if build_failed else InvokeStatus.tests_failed) From ff78deecaff624ab22a36bcc2273d9bfcb1f8d01 Mon Sep 17 00:00:00 2001 From: Maddipatla Chatan Date: Tue, 29 Sep 2026 07:14:05 +0530 Subject: [PATCH 26/58] Update ci/official/utilities/extract_resultstore_links.py Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- ci/official/utilities/extract_resultstore_links.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ci/official/utilities/extract_resultstore_links.py b/ci/official/utilities/extract_resultstore_links.py index f0e2a59bd6d13c..bd1fe562113cc1 100644 --- a/ci/official/utilities/extract_resultstore_links.py +++ b/ci/official/utilities/extract_resultstore_links.py @@ -115,7 +115,7 @@ def parse_log(file_path: str, tests_failed = re.search(TESTS_FAILED_RE, backtrack_line) if build_failed or tests_failed: log_fragment = '\n'.join( - log_lines[max(k - 20, 0):min(end_line + 1, len(log_lines))]) + log_lines[max(k - 20, 0):end_line + 1]) lines['log_fragment'] = log_fragment lines['status'] = (InvokeStatus.build_failed if build_failed else InvokeStatus.tests_failed) From 19cd995526a3169dc6fbdae60b3e9aeb815eddbc Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Mon, 28 Sep 2026 21:10:36 -0700 Subject: [PATCH 27/58] Address review: length check outside rank guard, allocation-free checks with NumPy error precedence, sub-2D tests --- .../python/ops/numpy_ops/np_array_ops.py | 30 ++++++++++++++----- .../python/ops/numpy_ops/tests/np_test.py | 9 ++++++ 2 files changed, 31 insertions(+), 8 deletions(-) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index e54a738066ecf5..cfe8f7cbc7f19f 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -1610,19 +1610,33 @@ def rot90(m, k=1, axes=(0, 1)): # pylint: disable=missing-docstring m = asarray(m) maybe_rank = m.shape.rank + if isinstance(axes, (tuple, list, range, np.ndarray)) and len(axes) != 2: + # Validate the sequence length even when the static rank is unknown: + # otherwise invalid lengths only surface as unpacking errors further + # down. + raise ValueError('len(axes) must be 2.') if maybe_rank is not None and isinstance( axes, (tuple, list, range, np.ndarray) ): - if len(axes) != 2: - raise ValueError('len(axes) must be 2.') - if builtins.all(isinstance(axis, (int, np.integer)) for axis in axes): - norm_axes = [axis + maybe_rank if axis < 0 else axis for axis in axes] - if builtins.any(axis < 0 or axis >= maybe_rank for axis in norm_axes): + # Direct scalar checks on the green path (no intermediate list, no + # generator expressions). NumPy checks for duplicate axes before + # out-of-range axes; `abs(ax0 - ax1) == maybe_rank` catches + # mixed-sign duplicates such as (0, -3) on a rank-3 array. + ax0, ax1 = axes[0], axes[1] + if isinstance(ax0, (int, np.integer)) and isinstance( + ax1, (int, np.integer) + ): + if ax0 == ax1 or builtins.abs(ax0 - ax1) == maybe_rank: + raise ValueError('Axes must be different.') + if ( + ax0 < -maybe_rank + or ax0 >= maybe_rank + or ax1 < -maybe_rank + or ax1 >= maybe_rank + ): raise ValueError( - f'Axes={tuple(axes)} out of range for array of rank {maybe_rank}.' + f'Axes={tuple(axes)} out of range for array of ndim={maybe_rank}.' ) - if norm_axes[0] == norm_axes[1]: - raise ValueError('Axes must be different.') m_rank = array_ops.rank(m) ax1, ax2 = np_utils._canonicalize_axes(axes, m_rank) # pylint: disable=protected-access diff --git a/tensorflow/python/ops/numpy_ops/tests/np_test.py b/tensorflow/python/ops/numpy_ops/tests/np_test.py index 352e849f1df191..f2397ff5341ff3 100644 --- a/tensorflow/python/ops/numpy_ops/tests/np_test.py +++ b/tensorflow/python/ops/numpy_ops/tests/np_test.py @@ -2092,6 +2092,15 @@ def testRot90InvalidAxes(self): # axes must have exactly two entries. with self.assertRaisesRegex(ValueError, 'must be 2'): tnp.rot90(a, axes=(0,)) + with self.assertRaisesRegex(ValueError, 'must be 2'): + tnp.rot90(a, axes=(0, 1, 2)) + # Sub-2D inputs are invalid with the default axes=(0, 1), matching + # NumPy's check order: the vector's default axes trip the duplicate + # check (abs(0 - 1) == ndim), and a scalar trips the bounds check. + with self.assertRaisesRegex(ValueError, 'out of range'): + tnp.rot90(tnp.ones(())) + with self.assertRaisesRegex(ValueError, 'must be different'): + tnp.rot90(tnp.ones(3)) # In-bounds negative axes remain valid. self.assertAllClose(tnp.rot90(a, axes=(-2, -1)), onp.rot90( onp.ones((2, 3)), axes=(-2, -1)), check_dtypes=False) From ded5b266100abf23d22aac37d0fcffc129a64a44 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Mon, 28 Sep 2026 21:11:06 -0700 Subject: [PATCH 28/58] Fix scalar transpose: convert axes with explicit int32 dtype so Tperm accepts the perm --- tensorflow/python/ops/numpy_ops/np_array_ops.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index b5e0a415758f3c..349dcca1ec68f4 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -997,7 +997,9 @@ def transpose(a, axes=None): raise ValueError('repeated axis in transpose') if axes is not None: - axes = asarray(axes) + # Specify an integer dtype explicitly: asarray([]) would otherwise + # infer float64, which the Transpose op's `Tperm` attr rejects. + axes = asarray(axes, dtype=np.int32) return array_ops.transpose(a=a, perm=axes) From b58dac4a45e8b38eece67f08a1ce0ec05b71bdb0 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Tue, 29 Sep 2026 08:06:39 -0700 Subject: [PATCH 29/58] Consolidate rot90 green-path checks under a single isinstance block --- .../python/ops/numpy_ops/np_array_ops.py | 45 +++++++++---------- 1 file changed, 22 insertions(+), 23 deletions(-) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index cfe8f7cbc7f19f..e63b3e9430032a 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -1610,33 +1610,32 @@ def rot90(m, k=1, axes=(0, 1)): # pylint: disable=missing-docstring m = asarray(m) maybe_rank = m.shape.rank - if isinstance(axes, (tuple, list, range, np.ndarray)) and len(axes) != 2: + if isinstance(axes, (tuple, list, range, np.ndarray)): # Validate the sequence length even when the static rank is unknown: # otherwise invalid lengths only surface as unpacking errors further # down. - raise ValueError('len(axes) must be 2.') - if maybe_rank is not None and isinstance( - axes, (tuple, list, range, np.ndarray) - ): - # Direct scalar checks on the green path (no intermediate list, no - # generator expressions). NumPy checks for duplicate axes before - # out-of-range axes; `abs(ax0 - ax1) == maybe_rank` catches - # mixed-sign duplicates such as (0, -3) on a rank-3 array. - ax0, ax1 = axes[0], axes[1] - if isinstance(ax0, (int, np.integer)) and isinstance( - ax1, (int, np.integer) - ): - if ax0 == ax1 or builtins.abs(ax0 - ax1) == maybe_rank: - raise ValueError('Axes must be different.') - if ( - ax0 < -maybe_rank - or ax0 >= maybe_rank - or ax1 < -maybe_rank - or ax1 >= maybe_rank + if len(axes) != 2: + raise ValueError('len(axes) must be 2.') + if maybe_rank is not None: + # Direct scalar checks on the green path (no intermediate list, no + # generator expressions). NumPy checks for duplicate axes before + # out-of-range axes; `abs(ax0 - ax1) == maybe_rank` catches + # mixed-sign duplicates such as (0, -3) on a rank-3 array. + ax0, ax1 = axes[0], axes[1] + if isinstance(ax0, (int, np.integer)) and isinstance( + ax1, (int, np.integer) ): - raise ValueError( - f'Axes={tuple(axes)} out of range for array of ndim={maybe_rank}.' - ) + if ax0 == ax1 or builtins.abs(ax0 - ax1) == maybe_rank: + raise ValueError('Axes must be different.') + if ( + ax0 < -maybe_rank + or ax0 >= maybe_rank + or ax1 < -maybe_rank + or ax1 >= maybe_rank + ): + raise ValueError( + f'Axes={tuple(axes)} out of range for array of ndim={maybe_rank}.' + ) m_rank = array_ops.rank(m) ax1, ax2 = np_utils._canonicalize_axes(axes, m_rank) # pylint: disable=protected-access From 9092f253a8c0236428d75e1b14faa30ac11a9c03 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Tue, 29 Sep 2026 08:07:09 -0700 Subject: [PATCH 30/58] Drop redundant post-loop duplicate check; in-loop bitmask rejection is sufficient --- tensorflow/python/ops/numpy_ops/np_array_ops.py | 5 ----- 1 file changed, 5 deletions(-) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index 349dcca1ec68f4..ad3cf15a0877b8 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -990,11 +990,6 @@ def transpose(a, axes=None): if normalized_mask & bit: raise ValueError('repeated axis in transpose') normalized_mask |= bit - seen_int_axes += 1 - if seen_int_axes == maybe_rank and normalized_mask != ( - 1 << maybe_rank - ) - 1: - raise ValueError('repeated axis in transpose') if axes is not None: # Specify an integer dtype explicitly: asarray([]) would otherwise From 63a87fb7b4674652b7e439a9d6f8b5d2b34309aa Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Tue, 29 Sep 2026 08:07:31 -0700 Subject: [PATCH 31/58] Remove unused seen_int_axes counter --- tensorflow/python/ops/numpy_ops/np_array_ops.py | 1 - 1 file changed, 1 deletion(-) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index ad3cf15a0877b8..9406a95862fb88 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -977,7 +977,6 @@ def transpose(a, axes=None): f"{len(axes)}." ) normalized_mask = 0 - seen_int_axes = 0 for ax in axes: if isinstance(ax, (int, np.integer)): normalized = ax + maybe_rank if ax < 0 else ax From 73cd721cb62eeccb6309ca459e6f623bc09ca4f2 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Tue, 29 Sep 2026 14:51:16 -0700 Subject: [PATCH 32/58] Convert rot90 axes to Python int before arithmetic to avoid unsigned modular wraparound --- tensorflow/python/ops/numpy_ops/np_array_ops.py | 5 +++++ tensorflow/python/ops/numpy_ops/tests/np_test.py | 10 ++++++++++ 2 files changed, 15 insertions(+) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index e63b3e9430032a..cb2a8321bfd062 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -1625,6 +1625,11 @@ def rot90(m, k=1, axes=(0, 1)): # pylint: disable=missing-docstring if isinstance(ax0, (int, np.integer)) and isinstance( ax1, (int, np.integer) ): + # Convert to Python int before arithmetic: NumPy unsigned scalars + # (e.g. np.uint32) perform modular subtraction, so `ax0 - ax1` + # would wrap around (e.g. 0 - 1 -> 4294967295) and both the + # duplicate check and the bounds check would misbehave. + ax0, ax1 = int(ax0), int(ax1) if ax0 == ax1 or builtins.abs(ax0 - ax1) == maybe_rank: raise ValueError('Axes must be different.') if ( diff --git a/tensorflow/python/ops/numpy_ops/tests/np_test.py b/tensorflow/python/ops/numpy_ops/tests/np_test.py index f2397ff5341ff3..afd3a559499ece 100644 --- a/tensorflow/python/ops/numpy_ops/tests/np_test.py +++ b/tensorflow/python/ops/numpy_ops/tests/np_test.py @@ -2101,9 +2101,19 @@ def testRot90InvalidAxes(self): tnp.rot90(tnp.ones(())) with self.assertRaisesRegex(ValueError, 'must be different'): tnp.rot90(tnp.ones(3)) + # Unsigned NumPy integers must not wrap around in the duplicate + # check (unsigned scalar subtraction is modular). + with self.assertRaisesRegex(ValueError, 'must be different'): + tnp.rot90(tnp.ones(3), axes=(np.uint32(0), np.uint32(0))) + with self.assertRaisesRegex(ValueError, 'must be different'): + tnp.rot90(tnp.ones(3), axes=(np.uint32(0), np.uint32(1))) # In-bounds negative axes remain valid. self.assertAllClose(tnp.rot90(a, axes=(-2, -1)), onp.rot90( onp.ones((2, 3)), axes=(-2, -1)), check_dtypes=False) + # Valid unsigned integer axes still work. + self.assertAllClose( + tnp.rot90(tnp.ones((2, 3)), axes=(np.uint32(0), np.uint32(1))), + onp.rot90(onp.ones((2, 3)), axes=(0, 1)), check_dtypes=False) # TODO(mattjj): test infix operator overrides From 8e5c13838c3ffc18457922a947d5ea85ab560e90 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Tue, 29 Sep 2026 19:37:47 -0700 Subject: [PATCH 33/58] Fix NameError in rot90 unsigned-axis tests: use onp.uint32 instead of np.uint32 np_test.py imports numpy as 'onp', so np.uint32 raised NameError at test collection. This matches the fix requested in the review of PR #127683. --- tensorflow/python/ops/numpy_ops/tests/np_test.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tensorflow/python/ops/numpy_ops/tests/np_test.py b/tensorflow/python/ops/numpy_ops/tests/np_test.py index afd3a559499ece..32cee1e589d8a0 100644 --- a/tensorflow/python/ops/numpy_ops/tests/np_test.py +++ b/tensorflow/python/ops/numpy_ops/tests/np_test.py @@ -2104,15 +2104,15 @@ def testRot90InvalidAxes(self): # Unsigned NumPy integers must not wrap around in the duplicate # check (unsigned scalar subtraction is modular). with self.assertRaisesRegex(ValueError, 'must be different'): - tnp.rot90(tnp.ones(3), axes=(np.uint32(0), np.uint32(0))) + tnp.rot90(tnp.ones(3), axes=(onp.uint32(0), onp.uint32(0))) with self.assertRaisesRegex(ValueError, 'must be different'): - tnp.rot90(tnp.ones(3), axes=(np.uint32(0), np.uint32(1))) + tnp.rot90(tnp.ones(3), axes=(onp.uint32(0), onp.uint32(1))) # In-bounds negative axes remain valid. self.assertAllClose(tnp.rot90(a, axes=(-2, -1)), onp.rot90( onp.ones((2, 3)), axes=(-2, -1)), check_dtypes=False) # Valid unsigned integer axes still work. self.assertAllClose( - tnp.rot90(tnp.ones((2, 3)), axes=(np.uint32(0), np.uint32(1))), + tnp.rot90(tnp.ones((2, 3)), axes=(onp.uint32(0), onp.uint32(1))), onp.rot90(onp.ones((2, 3)), axes=(0, 1)), check_dtypes=False) # TODO(mattjj): test infix operator overrides From 398deb9fc79104f32fda4f21dcee9bba88ecf718 Mon Sep 17 00:00:00 2001 From: kaivalya-cyber <141600539+kaivalya-cyber@users.noreply.github.com> Date: Tue, 29 Sep 2026 21:32:57 -0700 Subject: [PATCH 34/58] Use plain abs() instead of builtins.abs in rot90 duplicate check Style/consistency: abs is not shadowed anywhere in np_array_ops.py, so the explicit builtins.abs qualifier is unidiomatic. Per review of PR #127683. --- tensorflow/python/ops/numpy_ops/np_array_ops.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index 1f946b90539420..ca0f661d6aea37 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -1654,7 +1654,7 @@ def rot90(m, k=1, axes=(0, 1)): # pylint: disable=missing-docstring # would wrap around (e.g. 0 - 1 -> 4294967295) and both the # duplicate check and the bounds check would misbehave. ax0, ax1 = int(ax0), int(ax1) - if ax0 == ax1 or builtins.abs(ax0 - ax1) == maybe_rank: + if ax0 == ax1 or abs(ax0 - ax1) == maybe_rank: raise ValueError('Axes must be different.') if ( ax0 < -maybe_rank From 6dedc2ceb99e373a82bd7f63a170aaeee15be387 Mon Sep 17 00:00:00 2001 From: Henning Becker Date: Tue, 29 Sep 2026 23:28:21 -0700 Subject: [PATCH 35/58] Disable libnvjitlink disk caching via the -no-cache option. Unless `-no-cache` is passed to `nvJitLinkCreate`, `libnvJitLink` loads `libcuda.so`, calls `cuInit(0)`, and uses the CUDA driver's internal export table during `nvJitLinkAddData` to read and write the driver's on-disk JIT cache (`~/.nv/ComputeCache`). Aside from being redundant with XLA's own compilation caching and unnecessarily initializing the CUDA driver during compilation, concurrent access to `~/.nv/ComputeCache` across parallel processes or threads can cause the driver's `AddToCache` or `Get_Cache_Entry` callbacks to fail with `CUDA_ERROR_UNKNOWN` (999), which `libnvJitLink` treats as a fatal `NVJITLINK_ERROR_INTERNAL` (`ERROR 999: AddToCache`) even when PTX compilation succeeded. Pass `-no-cache` to `nvJitLinkCreate` in `CompileAndLinkUsingLibNvJitLink` and `GetLatestPtxIsaVersionForLibNvJitLink` in xla/stream_executor/cuda/nvjitlink.cc to disable the driver-level cache and work around these failures. PiperOrigin-RevId: 990769245 --- third_party/xla/xla/stream_executor/cuda/nvjitlink.cc | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/third_party/xla/xla/stream_executor/cuda/nvjitlink.cc b/third_party/xla/xla/stream_executor/cuda/nvjitlink.cc index 9b043eab05fa1f..be957aa107dca4 100644 --- a/third_party/xla/xla/stream_executor/cuda/nvjitlink.cc +++ b/third_party/xla/xla/stream_executor/cuda/nvjitlink.cc @@ -169,6 +169,7 @@ absl::StatusOr CompileAndLinkUsingLibNvJitLink( } cli_args.emplace_back("-Xptxas=--warn-on-spills"); cli_args.emplace_back(absl::StrCat("-split-compile=", inputs.size())); + cli_args.emplace_back("-no-cache"); if (options.disable_gpuasm_optimizations) { cli_args.emplace_back("-Xptxas=-O0"); @@ -284,7 +285,7 @@ absl::StatusOr GetLatestPtxIsaVersionForLibNvJitLink() { absl::string_view ptx_contents = ".version 99.99"; // The call to `nvJitLinkCreate` below requires an arch to be specified in // order to succeed. - std::vector cli_args_ptrs{"-arch=sm_90a"}; + std::vector cli_args_ptrs{"-arch=sm_90a", "-no-cache"}; nvJitLinkHandle link_handle = nullptr; nvJitLinkResult create_result = nvJitLinkCreate(&link_handle, /*num_args=*/cli_args_ptrs.size(), From b1e0513816ca9add749ba9c5e679ac45cc3c59fc Mon Sep 17 00:00:00 2001 From: Hyeontaek Lim Date: Tue, 29 Sep 2026 23:30:51 -0700 Subject: [PATCH 36/58] [IFRT Proxy] Avoid client error masking in client disconnection This change surfaces both client and server errors (if present). The client error often contains a meaningful error and is helpful for troubleshooting and testing. PiperOrigin-RevId: 990770193 --- .../xla/xla/python/ifrt_proxy/client/BUILD | 2 +- .../ifrt_proxy/client/grpc_client_session.cc | 21 +++++++++++++--- .../client/grpc_client_session_test.cc | 25 ++++++++++++++----- 3 files changed, 38 insertions(+), 10 deletions(-) diff --git a/third_party/xla/xla/python/ifrt_proxy/client/BUILD b/third_party/xla/xla/python/ifrt_proxy/client/BUILD index 3c78bf8293ab30..2d3dc283e8b30f 100644 --- a/third_party/xla/xla/python/ifrt_proxy/client/BUILD +++ b/third_party/xla/xla/python/ifrt_proxy/client/BUILD @@ -35,6 +35,7 @@ cc_library( "//xla/python/ifrt_proxy/common:grpc_ifrt_service_proto_cc", "//xla/python/ifrt_proxy/common:ifrt_service_proto_cc", "//xla/tsl/concurrency:future", + "//xla/tsl/platform:errors", "@com_google_absl//absl/base", "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/container:flat_hash_map", @@ -49,7 +50,6 @@ cc_library( "@grpc", "@grpc//:grpc++", "@tsl//tsl/platform:env", - "@tsl//tsl/platform:errors", "@tsl//tsl/platform:logging", "@tsl//tsl/platform:unbounded_work_queue", "@tsl//tsl/profiler/lib:traceme", diff --git a/third_party/xla/xla/python/ifrt_proxy/client/grpc_client_session.cc b/third_party/xla/xla/python/ifrt_proxy/client/grpc_client_session.cc index bbba5dd1872dd2..62e511a0c3ad9f 100644 --- a/third_party/xla/xla/python/ifrt_proxy/client/grpc_client_session.cc +++ b/third_party/xla/xla/python/ifrt_proxy/client/grpc_client_session.cc @@ -44,8 +44,8 @@ #include "xla/python/ifrt_proxy/common/grpc_ifrt_service.pb.h" #include "xla/python/ifrt_proxy/common/ifrt_service.pb.h" #include "xla/tsl/concurrency/future.h" +#include "xla/tsl/platform/errors.h" #include "tsl/platform/env.h" -#include "tsl/platform/errors.h" #include "tsl/platform/logging.h" #include "tsl/platform/threadpool.h" #include "tsl/platform/unbounded_work_queue.h" @@ -222,8 +222,23 @@ void GrpcClientSession::Finish(const absl::Status& client_status) { return server_status; }; - absl::Status combined_status = finish_stream_and_get_server_status(); - combined_status.Update(client_status); + // Prioritize `client_status` if non-OK because `context_->TryCancel()` + // above causes `stream_->Finish()` to return `CANCELLED` even when the + // stream was healthy before client-initiated termination. If both + // `client_status` and `server_status` are non-OK, include both in the + // combined error message so neither error is lost. + absl::Status server_status = finish_stream_and_get_server_status(); + absl::Status combined_status; + if (!client_status.ok() && !server_status.ok()) { + combined_status = tsl::errors::CreateWithUpdatedMessage( + client_status, + absl::StrCat("Client error: ", client_status.ToString(), + "; Server error: ", server_status.ToString())); + } else if (!client_status.ok()) { + combined_status = client_status; + } else { + combined_status = server_status; + } auto all_callbacks = response_callbacks_->PopAll(); for (auto& [_, cb] : all_callbacks) { diff --git a/third_party/xla/xla/python/ifrt_proxy/client/grpc_client_session_test.cc b/third_party/xla/xla/python/ifrt_proxy/client/grpc_client_session_test.cc index b53369ef72794d..dbb1967554907f 100644 --- a/third_party/xla/xla/python/ifrt_proxy/client/grpc_client_session_test.cc +++ b/third_party/xla/xla/python/ifrt_proxy/client/grpc_client_session_test.cc @@ -63,6 +63,7 @@ namespace proxy { namespace { +using ::testing::HasSubstr; using ::testing::Not; // Sufficient time for all processing (that are not explicitly waiting for @@ -281,7 +282,9 @@ TEST(GrpcClientSessionTest, HappyCaseTwoRequestsWithClientFinish) { EXPECT_EQ(cs.client_finished_q()->PopOrTimeout(), std::nullopt); cs.client_session()->Finish(TestError()); - EXPECT_THAT(cs.client_finished_q()->Pop(), Not(absl_testing::IsOk())); + EXPECT_THAT(cs.client_finished_q()->Pop(), + absl_testing::StatusIs(TestError().code(), + HasSubstr(TestError().message()))); } TEST(GrpcClientSessionTest, ServerFinishesDuringFirstRead) { @@ -323,12 +326,16 @@ TEST(GrpcClientSessionTest, ClientFinishesAfterServerConsumesFirstRequest) { session_ptr.store(cs.client_session()); TF_ASSERT_OK_AND_ASSIGN(Queue * response_q_1, cs.SendSimpleRequest()); - EXPECT_THAT(response_q_1->Pop(), Not(absl_testing::IsOk())); + EXPECT_THAT(response_q_1->Pop(), + absl_testing::StatusIs(TestError().code(), + HasSubstr(TestError().message()))); absl::StatusOr response_q_2 = cs.SendSimpleRequest(); EXPECT_THAT(response_q_2.status(), Not(absl_testing::IsOk())); - EXPECT_THAT(cs.client_finished_q()->Pop(), Not(absl_testing::IsOk())); + EXPECT_THAT(cs.client_finished_q()->Pop(), + absl_testing::StatusIs(TestError().code(), + HasSubstr(TestError().message()))); } TEST(GrpcClientSessionTest, ClientFinishesAfterServerWritesFirstResponse) { @@ -355,10 +362,14 @@ TEST(GrpcClientSessionTest, ClientFinishesAfterServerWritesFirstResponse) { // enqueued. If it could be enqueued, the client will die without the server // sending the corresponding response. if (response_q_2.ok()) { - EXPECT_THAT(response_q_2.value()->Pop(), Not(absl_testing::IsOk())); + EXPECT_THAT(response_q_2.value()->Pop(), + absl_testing::StatusIs(TestError().code(), + HasSubstr(TestError().message()))); } - EXPECT_THAT(cs.client_finished_q()->Pop(), Not(absl_testing::IsOk())); + EXPECT_THAT(cs.client_finished_q()->Pop(), + absl_testing::StatusIs(TestError().code(), + HasSubstr(TestError().message()))); } TEST(GrpcClientSessionTest, ClientFinishesDuringServerConstruction) { @@ -386,7 +397,9 @@ TEST(GrpcClientSessionTest, ClientFinishesDuringServerConstruction) { ExpectHeadAndTail({response_q_1, response_q_2}); - EXPECT_THAT(cs.client_finished_q()->Pop(), Not(absl_testing::IsOk())); + EXPECT_THAT(cs.client_finished_q()->Pop(), + absl_testing::StatusIs(TestError().code(), + HasSubstr(TestError().message()))); } TEST(GrpcClientSessionTest, MethodsAfterFinishReturnError) { From 4b89a6131d6579346d89d46bbc002f42f286c839 Mon Sep 17 00:00:00 2001 From: Mikhail Goncharov Date: Wed, 30 Sep 2026 00:05:34 -0700 Subject: [PATCH 37/58] [XLA:GPU] fix number of symbols when simplifying tiles number of symbols is dims + runtime vars, it's funny that we have not seen that fired before. Initialize range_var_indexing to tile_sizes to make it more clear that those are ts0, ts1,.. symbols. That is stills no-op but more readable. PiperOrigin-RevId: 990786285 --- .../tiling/experimental/tiling_space.cc | 33 +++++++++++-------- 1 file changed, 20 insertions(+), 13 deletions(-) diff --git a/third_party/xla/xla/codegen/tiling/experimental/tiling_space.cc b/third_party/xla/xla/codegen/tiling/experimental/tiling_space.cc index 2cfe2bb9b0521a..75887b6dc4345f 100644 --- a/third_party/xla/xla/codegen/tiling/experimental/tiling_space.cc +++ b/third_party/xla/xla/codegen/tiling/experimental/tiling_space.cc @@ -633,19 +633,22 @@ int64_t TilingSpace::num_parallel_dimensions() const { void TilingSpace::InitSimplificationIndexing() { CHECK(!is_symbolic_) << "Tile sizes must be assigned before initializing " "cached indexing map variables."; + CHECK(dim_vars_indexing_.empty()) + << "InitSimplificationIndexing must be called once"; + CHECK(range_vars_indexing_.empty()); + CHECK(rt_vars_indexing_.empty()); - dim_vars_indexing_.clear(); dim_vars_indexing_.reserve(dimensions_.size()); - for (const auto& dim_info : dimensions_) { - CHECK_GT(dim_info.tile_size.value(), 0); - int64_t upper_bound = - llvm::divideCeil(dim_info.dimension_size, dim_info.tile_size.value()); + range_vars_indexing_.reserve(dimensions_.size()); + for (const DimensionInfo& dim_info : dimensions_) { + int64_t tile_size = dim_info.tile_size.value(); + CHECK_GT(tile_size, 0); + int64_t upper_bound = llvm::divideCeil(dim_info.dimension_size, tile_size); dim_vars_indexing_.push_back(IndexingMap::Variable{0, upper_bound - 1}); + // Even though ts_X must already be replaced with constants right now, we + // initialize their bounds to [tile_size, tile_size] for completeness. + range_vars_indexing_.push_back(IndexingMap::Variable{tile_size, tile_size}); } - - range_vars_indexing_.assign(dimensions_.size(), IndexingMap::Variable{0, 0}); - - rt_vars_indexing_.clear(); rt_vars_indexing_.reserve(rt_vars_.size()); for (const auto& rt_var : rt_vars_) { rt_vars_indexing_.push_back(IndexingMap::Variable{rt_var.bounds}); @@ -662,10 +665,14 @@ llvm::SmallVector TilingSpace::SimplifyExpressions( } return simplified_expressions; } - // TODO(b/565301234): add constraints from tiling space? - SymbolicMap map = SymbolicMap::Get(mlir_context(), dimensions_.size(), - rt_vars_.size(), expressions); - + CHECK_EQ(dimensions_.size(), dim_vars_indexing_.size()); + CHECK_EQ(dimensions_.size(), range_vars_indexing_.size()); + CHECK_EQ(rt_vars_indexing_.size(), rt_vars_.size()); + // TODO(b/565301234): add constraints from tiling space? They don't seem to + // be used in the current implementation. + SymbolicMap map = + SymbolicMap::Get(mlir_context(), dimensions_.size(), + dimensions_.size() + rt_vars_.size(), expressions); IndexingMap indexing_map(map, dim_vars_indexing_, range_vars_indexing_, rt_vars_indexing_); indexing_map.Simplify(IndexingMap::SimplifyPointDimensions::kPreserve); From 863f92aea1c2f48407471156f70575c4ae489f18 Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Wed, 30 Sep 2026 00:11:35 -0700 Subject: [PATCH 38/58] Automated Code Change PiperOrigin-RevId: 990789026 --- tensorflow/core/kernels/deserialize_sparse_string_op.cc | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tensorflow/core/kernels/deserialize_sparse_string_op.cc b/tensorflow/core/kernels/deserialize_sparse_string_op.cc index d05473d5a6e525..9856e59ae40c84 100644 --- a/tensorflow/core/kernels/deserialize_sparse_string_op.cc +++ b/tensorflow/core/kernels/deserialize_sparse_string_op.cc @@ -203,7 +203,7 @@ class DeserializeSparseOp : public OpKernel { target_shape.vec()(i) = serialized_sparse.shape().dim_size(i); } for (int i = 0; i < output.dims() - 1; ++i) { - target_shape.vec()(i + ndims - 1) = output.shape().data()[i + 1]; + target_shape.vec()(i + ndims - 1) = output.shape()[i + 1]; } ReshapeSparseTensor(context, output.indices(), input_shape, From 4e9a74354bd5b6df6fa58eba09612bfe4b20d2fb Mon Sep 17 00:00:00 2001 From: Sohaib Iftikhar Date: Wed, 30 Sep 2026 00:41:39 -0700 Subject: [PATCH 39/58] [XLA:GPU] Skip collective symmetric buffer tests on non-Hopper architectures. PiperOrigin-RevId: 990801444 --- .../backends/gpu/tests/collective_ops_e2e_test.cc | 14 ++++++++++++++ .../gpu/tests/collective_ops_e2e_test_base.h | 4 ++-- 2 files changed, 16 insertions(+), 2 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test.cc b/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test.cc index f59c95a6dc0b29..0ad0026d5c032a 100644 --- a/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test.cc @@ -194,6 +194,13 @@ class AsyncCollectiveOps } protected: + void SetUp() override { + CollectiveOpsWithFlagsBase::SetUp(); + if (!IsHopperAndHigher() && enable_symmetric_buffer_) { + GTEST_SKIP() << "Test requires Hopper or higher"; + } + } + DebugOptions GetDebugOptionsForTest() const override { DebugOptions debug_options = CollectiveOpsWithFlagsBase::GetDebugOptionsForTest(); @@ -3846,6 +3853,13 @@ class SymmetricBufferCollectiveOpsTest : public CollectiveOpsTestE2E { : CollectiveOpsTestE2E(/*memory_size=*/128 * kMB, /*collectives_memory_size=*/64 * kMB) {} + void SetUp() override { + CollectiveOpsTestE2E::SetUp(); + if (!IsHopperAndHigher()) { + GTEST_SKIP() << "Test requires Hopper or higher"; + } + } + DebugOptions GetDebugOptionsForTest() const override { DebugOptions options = CollectiveOpsTestE2E::GetDebugOptionsForTest(); options.set_xla_gpu_enable_nccl_user_buffers(true); diff --git a/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test_base.h b/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test_base.h index 7bbd51b3213dbb..981495b9eb7c5e 100644 --- a/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test_base.h +++ b/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test_base.h @@ -83,12 +83,12 @@ class CollectiveOpsE2ETestBase : public gpu::HloPjRtGpuTestBase { return device_description().gpu_compute_capability(); } - bool IsHopperAndHigher() { + bool IsHopperAndHigher() const { return Capability().IsCuda() && Capability().cuda_compute_capability()->IsAtLeastHopper(); } - bool IsAmpereAndHigher() { + bool IsAmpereAndHigher() const { return Capability().IsCuda() && Capability().cuda_compute_capability()->IsAtLeastAmpere(); } From 29b9c9f4822aef267acd7b0074a8f0daedd38a90 Mon Sep 17 00:00:00 2001 From: Sohaib Iftikhar Date: Wed, 30 Sep 2026 00:41:42 -0700 Subject: [PATCH 40/58] [XLA:GPU] Add PerDeviceState container for per-device runtime state. Adds `xla::gpu::PerDeviceState`, which combines pre-allocated `VectorStorage` for device ordinals in `[0, num_devices)` with copy-on-write `CowStorage` fallback for ordinals outside `[0, num_devices)`. PiperOrigin-RevId: 990801463 --- .../xla/xla/backends/gpu/runtime/BUILD | 36 +++ .../xla/backends/gpu/runtime/device_slot.h | 5 + .../backends/gpu/runtime/per_device_state.cc | 99 ++++++++ .../backends/gpu/runtime/per_device_state.h | 105 ++++++++ .../gpu/runtime/per_device_state_test.cc | 239 ++++++++++++++++++ 5 files changed, 484 insertions(+) create mode 100644 third_party/xla/xla/backends/gpu/runtime/per_device_state.cc create mode 100644 third_party/xla/xla/backends/gpu/runtime/per_device_state.h create mode 100644 third_party/xla/xla/backends/gpu/runtime/per_device_state_test.cc diff --git a/third_party/xla/xla/backends/gpu/runtime/BUILD b/third_party/xla/xla/backends/gpu/runtime/BUILD index f5c3acd8945ea8..67bd53871d044c 100644 --- a/third_party/xla/xla/backends/gpu/runtime/BUILD +++ b/third_party/xla/xla/backends/gpu/runtime/BUILD @@ -5438,8 +5438,10 @@ cc_library( hdrs = ["device_slot.h"], visibility = ["//visibility:private"], deps = [ + "@com_google_absl//absl/base", "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/functional:function_ref", + "@com_google_absl//absl/status", ], ) @@ -5491,6 +5493,40 @@ xla_cc_test( ], ) +cc_library( + name = "per_device_state", + srcs = ["per_device_state.cc"], + hdrs = ["per_device_state.h"], + visibility = ["//visibility:private"], + deps = [ + ":cow_storage", + ":device_slot", + ":vector_storage", + "@com_google_absl//absl/base", + "@com_google_absl//absl/functional:function_ref", + "@com_google_absl//absl/status", + "@com_google_absl//absl/status:status_macros", + "@com_google_absl//absl/status:statusor", + "@com_google_absl//absl/strings", + ], +) + +xla_cc_test( + name = "per_device_state_test", + srcs = ["per_device_state_test.cc"], + deps = [ + ":per_device_state", + "//xla/tsl/platform:env", + "@com_google_absl//absl/base:core_headers", + "@com_google_absl//absl/container:flat_hash_set", + "@com_google_absl//absl/status", + "@com_google_absl//absl/status:status_matchers", + "@com_google_absl//absl/status:statusor", + "@com_google_absl//absl/synchronization", + "@com_google_googletest//:gtest_main", + ], +) + cc_library( name = "custom_kernel_thunk", srcs = ["custom_kernel_thunk.cc"], diff --git a/third_party/xla/xla/backends/gpu/runtime/device_slot.h b/third_party/xla/xla/backends/gpu/runtime/device_slot.h index 02640185ab6f97..732dc6f5b76c9d 100644 --- a/third_party/xla/xla/backends/gpu/runtime/device_slot.h +++ b/third_party/xla/xla/backends/gpu/runtime/device_slot.h @@ -18,8 +18,10 @@ limitations under the License. #include +#include "absl/base/call_once.h" #include "absl/base/optimization.h" #include "absl/functional/function_ref.h" +#include "absl/status/status.h" namespace xla::gpu { @@ -33,6 +35,9 @@ struct alignas(ABSL_CACHELINE_SIZE) DeviceSlot { template static T* Unwrap(DeviceSlot* slot); + + absl::once_flag init_flag; + absl::Status init_status; }; template diff --git a/third_party/xla/xla/backends/gpu/runtime/per_device_state.cc b/third_party/xla/xla/backends/gpu/runtime/per_device_state.cc new file mode 100644 index 00000000000000..009007fccc8eda --- /dev/null +++ b/third_party/xla/xla/backends/gpu/runtime/per_device_state.cc @@ -0,0 +1,99 @@ +/* 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/backends/gpu/runtime/per_device_state.h" + +#include +#include + +#include "absl/status/status.h" +#include "absl/status/statusor.h" +#include "absl/strings/str_cat.h" +#include "xla/backends/gpu/runtime/cow_storage.h" +#include "xla/backends/gpu/runtime/device_slot.h" +#include "xla/backends/gpu/runtime/vector_storage.h" + +namespace xla::gpu { +namespace { + +std::vector> CreateSlots( + int num_devices, DeviceSlotFactoryRef factory) { + if (num_devices <= 0) { + return {}; + } + std::vector> slots; + slots.reserve(num_devices); + for (int i = 0; i < num_devices; ++i) { + slots.push_back(factory()); + } + return slots; +} + +} // namespace + +class UntypedPerDeviceState::Impl { + public: + Impl(int num_devices, DeviceSlotFactoryRef factory) + : vector_storage_(CreateSlots(num_devices, factory)) {} + + int num_device_slots() const { return vector_storage_.size(); } + + DeviceSlot* Find(int device_ordinal) const { + if (device_ordinal < 0) { + return nullptr; + } + if (DeviceSlot* slot = vector_storage_.Find(device_ordinal)) { + return slot; + } + return cow_storage_.Find(device_ordinal); + } + + absl::StatusOr GetOrCreate(int device_ordinal, + DeviceSlotFactoryRef factory) { + if (device_ordinal < 0) { + return absl::InvalidArgumentError( + absl::StrCat("Negative device ordinal: ", device_ordinal)); + } + if (DeviceSlot* slot = vector_storage_.Find(device_ordinal)) { + return slot; + } + return cow_storage_.GetOrCreate(device_ordinal, factory); + } + + private: + VectorStorage vector_storage_; + CowStorage cow_storage_; +}; + +UntypedPerDeviceState::UntypedPerDeviceState(int num_devices, + DeviceSlotFactoryRef factory) + : impl_(std::make_unique(num_devices, factory)) {} + +UntypedPerDeviceState::~UntypedPerDeviceState() = default; + +int UntypedPerDeviceState::num_device_slots() const { + return impl_->num_device_slots(); +} + +DeviceSlot* UntypedPerDeviceState::Find(int device_ordinal) const { + return impl_->Find(device_ordinal); +} + +absl::StatusOr UntypedPerDeviceState::GetOrCreate( + int device_ordinal, DeviceSlotFactoryRef factory) { + return impl_->GetOrCreate(device_ordinal, factory); +} + +} // namespace xla::gpu diff --git a/third_party/xla/xla/backends/gpu/runtime/per_device_state.h b/third_party/xla/xla/backends/gpu/runtime/per_device_state.h new file mode 100644 index 00000000000000..7fef1fb3010b16 --- /dev/null +++ b/third_party/xla/xla/backends/gpu/runtime/per_device_state.h @@ -0,0 +1,105 @@ +/* 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. +==============================================================================*/ + +#ifndef XLA_BACKENDS_GPU_RUNTIME_PER_DEVICE_STATE_H_ +#define XLA_BACKENDS_GPU_RUNTIME_PER_DEVICE_STATE_H_ + +#include + +#include "absl/base/call_once.h" +#include "absl/functional/function_ref.h" +#include "absl/status/status.h" +#include "absl/status/status_macros.h" +#include "absl/status/statusor.h" +#include "xla/backends/gpu/runtime/device_slot.h" + +namespace xla::gpu { + +// Type-erased backing container for `PerDeviceState`. +class UntypedPerDeviceState { + public: + UntypedPerDeviceState(int num_devices, DeviceSlotFactoryRef factory); + ~UntypedPerDeviceState(); + + UntypedPerDeviceState(const UntypedPerDeviceState&) = delete; + UntypedPerDeviceState& operator=(const UntypedPerDeviceState&) = delete; + + int num_device_slots() const; + + DeviceSlot* Find(int device_ordinal) const; + absl::StatusOr GetOrCreate(int device_ordinal, + DeviceSlotFactoryRef factory); + + private: + class Impl; + std::unique_ptr impl_; +}; + +// Per-device states of type `T`, keyed by device ordinal. +// +// States for ordinals in [0, num_devices) are built in the constructor, and +// `Find` is an array index. States for any other non-negative ordinal are +// built by the first `GetOrCreate`, and `Find` is a lock-free scan over the +// ordinals built so far. `num_devices <= 0` (no topology) builds every state +// on first touch. +// +// Every state has its own cache line and keeps its address for the lifetime of +// the storage. `T` must be default-constructible. It is never copied or moved. +// The storage synchronizes only construction and publication across devices; +// callers must not mutate the same ordinal's `T` concurrently unless `T` is +// internally synchronized. +template +class PerDeviceState : private UntypedPerDeviceState { + public: + PerDeviceState() : PerDeviceState(0) {} + explicit PerDeviceState(int num_devices) + : UntypedPerDeviceState(num_devices, &DeviceSlot::Create) {} + + using UntypedPerDeviceState::num_device_slots; + + // Lock-free. Returns nullptr for an ordinal outside [0, num_devices) whose + // state was not built yet. + T* Find(int device_ordinal) const { + return DeviceSlot::Unwrap(UntypedPerDeviceState::Find(device_ordinal)); + } + + // Returns the state for `device_ordinal`, building it on the first call if + // the ordinal is outside [0, num_devices). Lock-free once the state exists. + absl::StatusOr GetOrCreate(int device_ordinal) { + ABSL_ASSIGN_OR_RETURN(DeviceSlot * slot, + UntypedPerDeviceState::GetOrCreate( + device_ordinal, &DeviceSlot::Create)); + return DeviceSlot::Unwrap(slot); + } + + // Thread-safe. Calls `init_fn` exactly once per `device_ordinal` and returns + // its status; subsequent calls for the same ordinal return the cached status. + // If concurrent callers pass different `init_fn`s for the same ordinal, it is + // unspecified which one is called. + absl::Status GetOrCreateAndInitialize( + int device_ordinal, absl::FunctionRef init_fn) { + ABSL_ASSIGN_OR_RETURN(DeviceSlot * slot, + UntypedPerDeviceState::GetOrCreate( + device_ordinal, &DeviceSlot::Create)); + absl::call_once(slot->init_flag, [&]() { + slot->init_status = init_fn(DeviceSlot::Unwrap(slot)); + }); + return slot->init_status; + } +}; + +} // namespace xla::gpu + +#endif // XLA_BACKENDS_GPU_RUNTIME_PER_DEVICE_STATE_H_ diff --git a/third_party/xla/xla/backends/gpu/runtime/per_device_state_test.cc b/third_party/xla/xla/backends/gpu/runtime/per_device_state_test.cc new file mode 100644 index 00000000000000..448e2794d38543 --- /dev/null +++ b/third_party/xla/xla/backends/gpu/runtime/per_device_state_test.cc @@ -0,0 +1,239 @@ +/* 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/backends/gpu/runtime/per_device_state.h" + +#include +#include +#include +#include +#include + +#include +#include +#include "absl/base/optimization.h" +#include "absl/container/flat_hash_set.h" +#include "absl/status/status.h" +#include "absl/status/status_matchers.h" +#include "absl/status/statusor.h" +#include "absl/synchronization/notification.h" +#include "xla/tsl/platform/env.h" +#include "xla/tsl/platform/threadpool.h" + +namespace xla::gpu { +namespace { + +using absl_testing::IsOkAndHolds; +using absl_testing::StatusIs; +using ::testing::Contains; +using ::testing::Not; + +// Counts constructions and destructions. Neither copyable nor movable, so a +// storage that copies or moves its states fails to compile. +struct Tracked { + static constexpr int kInitialValue = 42; + + Tracked() { constructed.fetch_add(1); } + ~Tracked() { destroyed.fetch_add(1); } + Tracked(const Tracked&) = delete; + Tracked& operator=(const Tracked&) = delete; + + static void Reset() { + constructed.store(0); + destroyed.store(0); + } + + static inline std::atomic constructed{0}; + static inline std::atomic destroyed{0}; + + int value = kInitialValue; +}; +static_assert(!std::is_copy_constructible_v); +static_assert(!std::is_copy_assignable_v); +static_assert(!std::is_move_constructible_v); +static_assert(!std::is_move_assignable_v); + +template +size_t CountDistinct(const std::vector& ptrs) { + return absl::flat_hash_set(ptrs.begin(), ptrs.end()).size(); +} + +template +size_t CountDistinctCacheLines(const std::vector& ptrs) { + absl::flat_hash_set lines; + for (const T* ptr : ptrs) { + lines.insert(reinterpret_cast(ptr) / ABSL_CACHELINE_SIZE); + } + return lines.size(); +} + +class PerDeviceStateTest : public ::testing::Test { + protected: + void SetUp() override { Tracked::Reset(); } +}; + +TEST_F(PerDeviceStateTest, OrdinalsBelowNumDevicesExistAfterConstruction) { + PerDeviceState storage(4); + EXPECT_EQ(storage.num_device_slots(), 4); + EXPECT_EQ(Tracked::constructed.load(), 4); + for (int ordinal = 0; ordinal < 4; ++ordinal) { + Tracked* state = storage.Find(ordinal); + EXPECT_NE(state, nullptr) << "ordinal " << ordinal; + EXPECT_THAT(storage.GetOrCreate(ordinal), IsOkAndHolds(state)) + << "ordinal " << ordinal; + } + EXPECT_EQ(Tracked::constructed.load(), 4); +} + +TEST_F(PerDeviceStateTest, WithoutDevicesStatesAreBuiltOnFirstTouch) { + for (int num_devices : {0, -1}) { + SCOPED_TRACE(::testing::Message() << "num_devices " << num_devices); + Tracked::Reset(); + PerDeviceState storage(num_devices); + EXPECT_EQ(storage.num_device_slots(), 0); + EXPECT_EQ(Tracked::constructed.load(), 0); + EXPECT_EQ(storage.Find(0), nullptr); + ASSERT_OK_AND_ASSIGN(Tracked * state, storage.GetOrCreate(0)); + EXPECT_NE(state, nullptr); + EXPECT_EQ(storage.Find(0), state); + EXPECT_EQ(Tracked::constructed.load(), 1); + } +} + +TEST_F(PerDeviceStateTest, OrdinalAtOrAboveNumDevicesIsBuiltOnFirstTouch) { + PerDeviceState storage(2); + EXPECT_EQ(storage.Find(2), nullptr); + EXPECT_EQ(storage.Find(5), nullptr); + ASSERT_OK_AND_ASSIGN(Tracked * state, storage.GetOrCreate(5)); + EXPECT_NE(state, nullptr); + EXPECT_EQ(storage.Find(5), state); + EXPECT_EQ(storage.Find(2), nullptr); + EXPECT_EQ(Tracked::constructed.load(), 3); +} + +TEST_F(PerDeviceStateTest, NegativeOrdinalIsRejected) { + PerDeviceState storage(2); + EXPECT_EQ(storage.Find(-1), nullptr); + EXPECT_THAT(storage.GetOrCreate(-1), + StatusIs(absl::StatusCode::kInvalidArgument)); + EXPECT_EQ(Tracked::constructed.load(), 2); +} + +TEST_F(PerDeviceStateTest, StatesAreDistinctAcrossStorages) { + // Ordinals 0 and 1 are built in the constructor, 2 and 3 on first touch. + PerDeviceState storage(2); + std::vector states; + for (int ordinal = 0; ordinal < 4; ++ordinal) { + ASSERT_OK_AND_ASSIGN(Tracked * state, storage.GetOrCreate(ordinal)); + states.push_back(state); + } + EXPECT_THAT(states, Not(Contains(nullptr))); + EXPECT_EQ(CountDistinct(states), 4); + EXPECT_EQ(CountDistinctCacheLines(states), 4); + for (int ordinal = 0; ordinal < 4; ++ordinal) { + EXPECT_EQ(storage.Find(ordinal), states[ordinal]) << "ordinal " << ordinal; + } + EXPECT_EQ(Tracked::constructed.load(), 4); +} + +TEST_F(PerDeviceStateTest, ConcurrentGetOrCreateAcrossStorages) { + constexpr int kNumDevices = 4; + constexpr int kNumOrdinals = 8; + constexpr int kThreadsPerOrdinal = 4; + constexpr int kNumThreads = kNumOrdinals * kThreadsPerOrdinal; + PerDeviceState storage(kNumDevices); + std::vector> results(kNumThreads); + absl::Notification start; + { + tsl::thread::ThreadPool pool(tsl::Env::Default(), "per_device_state_test", + kNumThreads); + for (int i = 0; i < kNumThreads; ++i) { + pool.Schedule([&, i] { + start.WaitForNotification(); + results[i] = storage.GetOrCreate(i % kNumOrdinals); + }); + } + start.Notify(); + } // Joins the pool. + EXPECT_EQ(Tracked::constructed.load(), kNumOrdinals); + for (int i = 0; i < kNumThreads; ++i) { + Tracked* state = storage.Find(i % kNumOrdinals); + EXPECT_NE(state, nullptr) << "thread " << i; + EXPECT_THAT(results[i], IsOkAndHolds(state)) << "thread " << i; + } +} + +TEST_F(PerDeviceStateTest, GetOrCreateAndInitializeRunsOncePerOrdinal) { + constexpr int kNumDevices = 2; + constexpr int kNumOrdinals = 4; + constexpr int kThreadsPerOrdinal = 4; + constexpr int kNumThreads = kNumOrdinals * kThreadsPerOrdinal; + PerDeviceState storage(kNumDevices); + std::atomic init_calls{0}; + std::vector statuses(kNumThreads); + absl::Notification start; + { + tsl::thread::ThreadPool pool(tsl::Env::Default(), "per_device_state_test", + kNumThreads); + for (int i = 0; i < kNumThreads; ++i) { + pool.Schedule([&, i] { + start.WaitForNotification(); + int ordinal = i % kNumOrdinals; + statuses[i] = + storage.GetOrCreateAndInitialize(ordinal, [&](Tracked* state) { + init_calls.fetch_add(1); + state->value = 100 + ordinal; + return absl::OkStatus(); + }); + }); + } + start.Notify(); + } // Joins the pool. + EXPECT_EQ(init_calls.load(), kNumOrdinals); + for (int i = 0; i < kNumThreads; ++i) { + EXPECT_THAT(statuses[i], absl_testing::IsOk()) << "thread " << i; + } + for (int ordinal = 0; ordinal < kNumOrdinals; ++ordinal) { + Tracked* state = storage.Find(ordinal); + ASSERT_NE(state, nullptr) << "ordinal " << ordinal; + EXPECT_EQ(state->value, 100 + ordinal) << "ordinal " << ordinal; + } +} + +TEST_F(PerDeviceStateTest, GetOrCreateAndInitializeCachesErrorStatus) { + PerDeviceState storage(2); + int init_calls = 0; + auto failing_init = [&](Tracked*) { + ++init_calls; + return absl::InternalError("init failed"); + }; + + EXPECT_THAT(storage.GetOrCreateAndInitialize(0, failing_init), + StatusIs(absl::StatusCode::kInternal)); + EXPECT_THAT(storage.GetOrCreateAndInitialize(0, + [&](Tracked*) { + ++init_calls; + return absl::OkStatus(); + }), + StatusIs(absl::StatusCode::kInternal)); + EXPECT_EQ(init_calls, 1); + + EXPECT_THAT(storage.GetOrCreateAndInitialize(-1, failing_init), + StatusIs(absl::StatusCode::kInvalidArgument)); + EXPECT_EQ(init_calls, 1); +} + +} // namespace +} // namespace xla::gpu From 40bfb8d760f740bacf6a8166dec39814c7ba2f98 Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Wed, 30 Sep 2026 01:24:53 -0700 Subject: [PATCH 41/58] Automated Code Change PiperOrigin-RevId: 990820587 --- third_party/xla/xla/hlo/evaluator/hlo_evaluator.cc | 2 +- third_party/xla/xla/hlo/evaluator/hlo_evaluator_test.cc | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/third_party/xla/xla/hlo/evaluator/hlo_evaluator.cc b/third_party/xla/xla/hlo/evaluator/hlo_evaluator.cc index 5637a5f897d5f6..ca6168ffcff37c 100644 --- a/third_party/xla/xla/hlo/evaluator/hlo_evaluator.cc +++ b/third_party/xla/xla/hlo/evaluator/hlo_evaluator.cc @@ -212,7 +212,7 @@ absl::Status MakeEvalErrorDueToParamOrInfeed( DCHECK(absl::endian::native == absl::endian::big); error_detail = absl::byteswap(error_detail); } - (*error_payload.data()) = error_detail; + (error_payload[0]) = error_detail; error.SetPayload(internal::kEvalErrorDetailUrl, absl::Cord(error_payload)); return error; } diff --git a/third_party/xla/xla/hlo/evaluator/hlo_evaluator_test.cc b/third_party/xla/xla/hlo/evaluator/hlo_evaluator_test.cc index e8f47905b40997..b4cf7b931c3a03 100644 --- a/third_party/xla/xla/hlo/evaluator/hlo_evaluator_test.cc +++ b/third_party/xla/xla/hlo/evaluator/hlo_evaluator_test.cc @@ -8844,7 +8844,7 @@ TEST(EvalErrorTest, Payload) { DCHECK(absl::endian::native == absl::endian::big); error_detail = absl::byteswap(error_detail); } - (*payload.data()) = error_detail; + (payload[0]) = error_detail; s.SetPayload(internal::kEvalErrorDetailUrl, absl::Cord(payload)); From c3f724048f5d98834bb96aa3fa7b84d4bbde90e3 Mon Sep 17 00:00:00 2001 From: Dirk Hornung Date: Wed, 30 Sep 2026 01:33:49 -0700 Subject: [PATCH 42/58] [XLA:GPU] Run the cuDNN frontend graph warmup on a private stream to avoid race conditions when sharing the stream between multiple threads. PiperOrigin-RevId: 990824347 --- .../xla/xla/stream_executor/cuda/BUILD | 1 + .../xla/xla/stream_executor/cuda/cuda_dnn.cc | 49 ++++++++++++++----- 2 files changed, 37 insertions(+), 13 deletions(-) diff --git a/third_party/xla/xla/stream_executor/cuda/BUILD b/third_party/xla/xla/stream_executor/cuda/BUILD index 5d69ba482808a8..b1a74c16f283d0 100644 --- a/third_party/xla/xla/stream_executor/cuda/BUILD +++ b/third_party/xla/xla/stream_executor/cuda/BUILD @@ -476,6 +476,7 @@ cc_library( ":cuda_compute_capability", ":cuda_diagnostics", ":cuda_platform_id", + ":cuda_status", ":cudnn_api_wrappers", ":cudnn_frontend_helpers", ":cudnn_sdpa_score_mod", diff --git a/third_party/xla/xla/stream_executor/cuda/cuda_dnn.cc b/third_party/xla/xla/stream_executor/cuda/cuda_dnn.cc index b4321feb17adfe..2839cd253d46d7 100644 --- a/third_party/xla/xla/stream_executor/cuda/cuda_dnn.cc +++ b/third_party/xla/xla/stream_executor/cuda/cuda_dnn.cc @@ -57,6 +57,7 @@ limitations under the License. #include "xla/stream_executor/cuda/cuda_compute_capability.h" #include "xla/stream_executor/cuda/cuda_diagnostics.h" #include "xla/stream_executor/cuda/cuda_platform_id.h" +#include "xla/stream_executor/cuda/cuda_status.h" #include "xla/stream_executor/cuda/cudnn_api_wrappers.h" #include "xla/stream_executor/cuda/cudnn_frontend_helpers.h" #include "xla/stream_executor/cuda/cudnn_sdpa_score_mod.h" @@ -213,11 +214,6 @@ class CudnnHandle { lock_(std::move(lock)), handle_(handle) {} - // Takes ownership of the lock to access cuDNN using handle. Doesn't activate - // a CUDA context. - CudnnHandle(std::unique_ptr lock, cudnnHandle_t handle) - : lock_(std::move(lock)), handle_(handle) {} - // Returns cuDNN handle. To be passed directly to cuDNN APIs, don't keep // a copy. cudnnHandle_t handle() const { return handle_; } @@ -269,6 +265,9 @@ class CudnnAccess { if (compilation_handle_) { cudnnDestroy(compilation_handle_); } + if (private_stream_) { + cuStreamDestroy(private_stream_); + } } // Creates a CudnnHandle instance for stream. @@ -302,14 +301,29 @@ class CudnnAccess { } // Creates a CudnnHandle instance for the compilation handle, which is used - // to build cuDNN graphs and execution plans. - absl::StatusOr GetCompilationHandle() { + // to build and deserialize cuDNN graphs and execution plans. + // + // TOOD(b/567795551): We currently use a private stream for all compilation + // related tasks to avoid race conditions. Explore if we can use the default + // stream instead. + absl::StatusOr GetCompilationHandle(StreamExecutor* executor) { auto lock = std::make_unique(compilation_mutex_); compilation_mutex_.AssertHeld(); if (!compilation_handle_) { return absl::InternalError("CudnnAccess not properly initialized."); } - return CudnnHandle(std::move(lock), compilation_handle_); + CudnnHandle cudnn(executor, std::move(lock), compilation_handle_); + if (private_stream_ == nullptr) { + CUstream stream; + ABSL_RETURN_IF_ERROR( + cuda::ToStatus(cuStreamCreate(&stream, CU_STREAM_NON_BLOCKING))); + if (cudnnSetStream(compilation_handle_, stream) != CUDNN_STATUS_SUCCESS) { + cuStreamDestroy(stream); + return absl::InternalError("Failed to set cuDNN compilation stream."); + } + private_stream_ = stream; + } + return cudnn; } void NotifyStreamDestroyed(Stream* stream) { @@ -340,6 +354,10 @@ class CudnnAccess { // Shared compilation handle for all threads calling GetCompilationHandle(). cudnnHandle_t compilation_handle_ ABSL_GUARDED_BY(compilation_mutex_) = nullptr; // Owned. + + // Private stream bound to compilation_handle_, see GetCompilationHandle(). + CUstream private_stream_ ABSL_GUARDED_BY(compilation_mutex_) = + nullptr; // Owned. }; namespace { @@ -6821,7 +6839,10 @@ bool CudnnSupport::DeriveOutputBatchDescriptor( absl::StatusOr> CudnnSupport::DeserializeGraph( Stream& stream, absl::string_view serialized_data) const { - auto cudnn = cudnn_->GetHandle(stream.parent(), &stream); + // Deserializing a graph runs a warmup execution by default. Use a private + // stream for the warmup to avoid race conditions. + ABSL_ASSIGN_OR_RETURN(CudnnHandle cudnn, + cudnn_->GetCompilationHandle(stream.parent())); cudnn_frontend::graph::Graph graph; RETURN_IF_CUDNN_FRONTEND_ERROR(graph.deserialize( cudnn.handle(), @@ -6904,8 +6925,9 @@ absl::Status CudnnGraph::Prepare(dnn::DnnSupport* dnn_support, const CudnnSupport& cudnn_support = static_cast(*dnn_support); // Holds the lock on the shared compilation handle until the end of scope. - ABSL_ASSIGN_OR_RETURN(CudnnHandle cudnn, - cudnn_support.cudnn_->GetCompilationHandle()); + ABSL_ASSIGN_OR_RETURN( + CudnnHandle cudnn, + cudnn_support.cudnn_->GetCompilationHandle(cudnn_support.parent_)); RETURN_IF_CUDNN_FRONTEND_ERROR(graph_.validate()); RETURN_IF_CUDNN_FRONTEND_ERROR( graph_.build_operation_graph(cudnn.handle())); @@ -6931,8 +6953,9 @@ absl::Status CudnnGraph::Build(dnn::DnnSupport* dnn_support, if (dnn_support) { const CudnnSupport& cudnn_support = static_cast(*dnn_support); - ABSL_ASSIGN_OR_RETURN(CudnnHandle cudnn, - cudnn_support.cudnn_->GetCompilationHandle()); + ABSL_ASSIGN_OR_RETURN( + CudnnHandle cudnn, + cudnn_support.cudnn_->GetCompilationHandle(cudnn_support.parent_)); if (plan_id.has_value()) { RETURN_CUDNN_FRONTEND_STATUS( graph_.build_plan_at_index(cudnn.handle(), *plan_id)); From 2816dcad6a53a80108bc4c66e8c4baa88cb3f387 Mon Sep 17 00:00:00 2001 From: dlcompilers-infra-bot Date: Wed, 30 Sep 2026 01:50:06 -0700 Subject: [PATCH 43/58] PR #49717: Export control_predecessors() to python MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Imported from GitHub PR https://github.com/openxla/xla/pull/49717 📝 Summary of Changes Export control_predecessors() to python 🎯 Justification JAX has an experimental API to allow python to modify the Thunk scheduling. To build valid one, we need to know the control depedence. 🚀 Kind of Contribution ✨ New Feature, 🧪 Tests 📊 Benchmark (for Performance Improvements) No speed-up expected by this exposure. 🧪 Unit Tests: The new Python API have been tested. 🧪 Execution Tests: xla/python/xla_hlo_test.py::TestHloModule.testHloInstructionControlPredecessors Copybara import of the project: -- 9c42acd7cc52413d38d851e914f65af1f5f21135 by Frederic Bastien : Export control_predecessors() to python Merging this change closes #49717 PiperOrigin-RevId: 990831343 --- third_party/xla/xla/python/_hlo.pyi | 1 + third_party/xla/xla/python/hlo.cc | 10 ++++++++++ third_party/xla/xla/python/xla_hlo_test.py | 19 +++++++++++++++++++ 3 files changed, 30 insertions(+) diff --git a/third_party/xla/xla/python/_hlo.pyi b/third_party/xla/xla/python/_hlo.pyi index 5b1505f46a492e..fe4fcf435d2457 100644 --- a/third_party/xla/xla/python/_hlo.pyi +++ b/third_party/xla/xla/python/_hlo.pyi @@ -563,6 +563,7 @@ class HloInstruction: def shape(self) -> Shape: ... def users(self) -> list[HloInstruction]: ... def operands(self) -> list[HloInstruction]: ... + def control_predecessors(self) -> list[HloInstruction]: ... def async_wrapped_root(self) -> HloInstruction: ... def get_frontend_attribute(self, key: str) -> str | None: ... def set_frontend_attribute(self, key: str, value: str) -> None: ... diff --git a/third_party/xla/xla/python/hlo.cc b/third_party/xla/xla/python/hlo.cc index b2b0b8ef88909b..23ee13cc2c8570 100644 --- a/third_party/xla/xla/python/hlo.cc +++ b/third_party/xla/xla/python/hlo.cc @@ -754,6 +754,15 @@ NB_MODULE(_hlo, m) { } return operands; } + std::vector> control_predecessors() + const { + std::vector> predecessors; + for (const HloInstruction* predecessor : inst_->control_predecessors()) { + predecessors.push_back( + std::make_shared(predecessor, module_)); + } + return predecessors; + } const HloInstruction* inst() const { return inst_; } std::shared_ptr async_wrapped_root() const { @@ -824,6 +833,7 @@ NB_MODULE(_hlo, m) { .def_prop_ro("shape", &InstructionWrapper::shape) .def("users", &InstructionWrapper::users) .def("operands", &InstructionWrapper::operands) + .def("control_predecessors", &InstructionWrapper::control_predecessors) .def("async_wrapped_root", &InstructionWrapper::async_wrapped_root) .def("get_frontend_attribute", &InstructionWrapper::get_frontend_attribute, nb::arg("key")) diff --git a/third_party/xla/xla/python/xla_hlo_test.py b/third_party/xla/xla/python/xla_hlo_test.py index 848a31184bf7ec..247fa2db16a963 100644 --- a/third_party/xla/xla/python/xla_hlo_test.py +++ b/third_party/xla/xla/python/xla_hlo_test.py @@ -280,6 +280,25 @@ def testHloInstructionOperands(self): self.assertIsInstance(operand, _hlo.HloInstruction) logging.info("Operand of %s: %s", inst.name, operand.name) + @unittest.skipIf(cloud_tpu or pathways, "not implemented") + def testHloInstructionControlPredecessors(self): + module = _hlo.hlo_module_from_text(R""" +HloModule control_predecessors + +ENTRY %main { + %p0 = f32[] parameter(0) + %first = f32[] negate(%p0) + ROOT %second = f32[] abs(%p0), control-predecessors={%first} +} +""") + instructions = { + inst.name: inst for inst in module.computations()[0].instructions() + } + self.assertEmpty(instructions["first"].control_predecessors()) + predecessors = instructions["second"].control_predecessors() + self.assertLen(predecessors, 1) + self.assertEqual(predecessors[0].name, "first") + @unittest.skipIf(cloud_tpu or pathways, "not implemented") def testHloInstructionName(self): module = self.ExampleComputation() From bbdf033a99bdf8953382dd646c7c0571455a73cd Mon Sep 17 00:00:00 2001 From: Marco Minutoli Date: Wed, 30 Sep 2026 01:50:50 -0700 Subject: [PATCH 44/58] PR #49741: [ROCm] Remove dead ROCm version checks from GemmRewriter MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Imported from GitHub PR https://github.com/openxla/xla/pull/49741 ## Summary - Remove all `toolkit_version_ < {6,x,0}` runtime checks in `gemm_rewriter.cc` — they are dead code since `rocm_config.h.tpl` requires ROCm >= 7.1 at compile time (`#error` if `TF_ROCM_VERSION < 70100`). - Simplify the `toolkit_version_ >= {7,0,0}` guard on Swish matching to plain `if (is_rocm)`, since it is always true. - Remove the now-callerless `TurnF8DotWithUnsupportedOutputTypeIntoF32()` helper. - Clean up corresponding dead branches in `gemm_rewriter_fp8_test.cc` (ROCm < 6.0 skip, and five if/else blocks selecting CHECK patterns for ROCm < 6.2). Copybara import of the project: -- 996cd22d95e984681c58ad0187e75c1fb698c666 by Marco Minutoli : [ROCm] Remove dead ROCm version checks from GemmRewriter rocm_config.h.tpl requires ROCm >= 7.1 (#error if < 70100), making all runtime checks against ROCm 6.x unreachable. Remove the dead guards: - toolkit_version_ < {6,2,0} output-type workaround at FP8 dot rewrite - toolkit_version_ < {6,0,0} FP8 availability check (has_fp8_support remains) - toolkit_version_ < {6,2,0} output-type restriction in CreateF8CustomCall - toolkit_version_ >= {7,0,0} always-true guard on Swish matching (simplified) - TurnF8DotWithUnsupportedOutputTypeIntoF32() helper (no remaining callers) - Corresponding dead branches in gemm_rewriter_fp8_test.cc Merging this change closes #49741 PiperOrigin-RevId: 990831715 --- .../backends/gpu/transforms/gemm_rewriter.cc | 38 +------------- .../gpu/transforms/gemm_rewriter_fp8_test.cc | 50 +++---------------- 2 files changed, 9 insertions(+), 79 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/transforms/gemm_rewriter.cc b/third_party/xla/xla/backends/gpu/transforms/gemm_rewriter.cc index 14c5bd07fe259b..6d376b5b6586c1 100644 --- a/third_party/xla/xla/backends/gpu/transforms/gemm_rewriter.cc +++ b/third_party/xla/xla/backends/gpu/transforms/gemm_rewriter.cc @@ -732,13 +732,6 @@ class GemmRewriterVisitor : public DfsHloRewriteVisitor { const_cast(instr->operand(0)))) && (b = MatchFp8Param( const_cast(instr->operand(1))))) { - if (gpu_version_.IsRocm() && - toolkit_version_ < stream_executor::SemanticVersion{6, 2, 0} && - instr->shape().element_type() != F16 && - instr->shape().element_type() != F32) { - ABSL_ASSIGN_OR_RETURN(instr, - TurnF8DotWithUnsupportedOutputTypeIntoF32(instr)); - } ABSL_ASSIGN_OR_RETURN(bool created_call, CreateF8CustomCall(instr, gpu_backend_config, a.value(), b.value())); @@ -973,8 +966,7 @@ class GemmRewriterVisitor : public DfsHloRewriteVisitor { } const auto is_rocm = gpu_version_.IsRocm(); - if (is_rocm && - toolkit_version_ >= stream_executor::SemanticVersion{7, 0, 0}) { + if (is_rocm) { // Attempt to match approximate Swish activation (including grouped // matmul) // (https://flax.readthedocs.io/en/v0.5.3/_autosummary/flax.linen.swish.html), @@ -1265,11 +1257,6 @@ class GemmRewriterVisitor : public DfsHloRewriteVisitor { VLOG(1) << "FP8 Custom Calls require MI300, or later architectures."; return false; } - if (toolkit_version_ < stream_executor::SemanticVersion{6, 0, 0}) { - // FP8 GEMM kernels are only available with ROCm 6.0 and above - VLOG(1) << "FP8 Custom Calls require ROCm 6.0 or newer."; - return false; - } } PrimitiveType a_type = a.fp8_input->shape().element_type(); @@ -1404,15 +1391,6 @@ class GemmRewriterVisitor : public DfsHloRewriteVisitor { } } if (gpu_version_.IsRocm()) { - if (toolkit_version_ < stream_executor::SemanticVersion{6, 2, 0}) { - if (supported_d_types.find(d_type) == supported_d_types.end()) { - VLOG(1) << "Failed to rewrite " << instr->ToShortString() - << " into FP8 Custom Call. For ROCm version < 6.2, output " - "type must be BF16, F16 or F32, but got " - << PrimitiveType_Name(d_type); - return false; - } - } ABSL_ASSIGN_OR_RETURN(auto rocm_compute_capability, GetRocmComputeCapability(gpu_version_)); if (rocm_compute_capability.has_ocp_fp8_support()) { @@ -2719,20 +2697,6 @@ class GemmRewriterVisitor : public DfsHloRewriteVisitor { return gemm_config.rhs_layout.num_cols <= kMaxDimensionSize; } - // Turns an F8 dot with unsupported output type into an F8 dot with F32 - // output, and converting the F32 output to unsupported output types. - absl::StatusOr TurnF8DotWithUnsupportedOutputTypeIntoF32( - HloInstruction* instr) { - Shape output_f32_shape = instr->shape(); - output_f32_shape.set_element_type(F32); - HloInstruction* f32_dot = - instr->AddInstruction(instr->CloneWithNewShape(output_f32_shape)); - HloInstruction* convert = instr->AddInstruction( - HloInstruction::CreateConvert(instr->shape(), f32_dot)); - ABSL_RETURN_IF_ERROR(ReplaceInstruction(instr, convert)); - return f32_dot; - } - // Turns an F8 dot into an F16 dot, converting operands to F16 (or BF16) and // converting the output back to F8. absl::StatusOr TurnF8DotIntoF16Dot(HloInstruction* instr) { diff --git a/third_party/xla/xla/backends/gpu/transforms/gemm_rewriter_fp8_test.cc b/third_party/xla/xla/backends/gpu/transforms/gemm_rewriter_fp8_test.cc index 06fc23babe4da7..4f3902d38b042a 100644 --- a/third_party/xla/xla/backends/gpu/transforms/gemm_rewriter_fp8_test.cc +++ b/third_party/xla/xla/backends/gpu/transforms/gemm_rewriter_fp8_test.cc @@ -91,11 +91,6 @@ class ParameterizedFp8GemmRewriteTest GTEST_SKIP() << "FP8 is not supported on this GPU architecture."; } - if (IsRocm() && GetToolkitVersion() < se::SemanticVersion{6, 0, 0}) { - GTEST_SKIP() - << "F8 gemm rewrite is only supported in ROCm 6.0 and above."; - } - if (IsRocm() && !Capability().rocm_compute_capability()->has_fp8_support()) { GTEST_SKIP() @@ -330,15 +325,9 @@ TEST_F(ParameterizedFp8GemmRewriteTest, UnscaledABUnscaledDF8) { ; CHECK-NEXT: [[P1_TRANSPOSE:%[^ ]+]] = <>[16,32]{1,0} transpose([[P1]]), dimensions={1,0} ; CHECK-NEXT: [[C1:[^ ]+]] = f32[] constant(1) )"; - if (IsRocm() && GetToolkitVersion() < se::SemanticVersion{6, 2, 0}) { - checks.append( - R"(; CHECK-GCN-NEXT: [[OUT:%[^ ]+]] = (f32[16,16]{1,0}, s8[{{[0-9]+}}]{0}) custom-call([[P0]], [[P1_TRANSPOSE]], [[C1]], [[C1]]), -)"); - } else { - checks.append( - R"(; CHECK-NEXT: [[OUT:%[^ ]+]] = (<>[16,16]{1,0}, s8[{{[0-9]+}}]{0}) custom-call([[P0]], [[P1_TRANSPOSE]], [[C1]], [[C1]]), + checks.append( + R"(; CHECK-NEXT: [[OUT:%[^ ]+]] = (<>[16,16]{1,0}, s8[{{[0-9]+}}]{0}) custom-call([[P0]], [[P1_TRANSPOSE]], [[C1]], [[C1]]), )"); - } checks.append( R"(; CHECK: custom_call_target="__cublas$lt$matmul$f8", ; CHECK: backend_config={ @@ -1158,15 +1147,9 @@ TEST_F(ParameterizedFp8GemmRewriteTest, ; CHECK-NEXT: [[P3:%[^ ]+]] = bf16[] parameter(3) ; CHECK-NEXT: [[XS1:%[^ ]+]] = f32[] convert([[P3]]) )"; - if (IsRocm() && GetToolkitVersion() < se::SemanticVersion{6, 2, 0}) { - checks += - R"(; CHECK-GCN-NEXT: [[OUT:%[^ ]+]] = (f32[16,16]{1,0}, s8[{{[0-9]+}}]{0}) custom-call([[P0]], [[P1_TRANSPOSE]], [[XS]], [[XS1]]), -)"; - } else { - checks += R"(; CHECK-NEXT: [[B:%[^ ]+]] = bf16[16]{0} parameter(4) + checks += R"(; CHECK-NEXT: [[B:%[^ ]+]] = bf16[16]{0} parameter(4) ; CHECK-NEXT: [[OUT:%[^ ]+]] = (bf16[16,16]{1,0}, s8[{{[0-9]+}}]{0}) custom-call([[P0]], [[P1_TRANSPOSE]], [[XS]], [[XS1]], [[B]]), )"; - } checks += R"(; CHECK: custom_call_target="__cublas$lt$matmul$f8", ; CHECK: backend_config={ ; CHECK-DAG: "alpha_real":1 @@ -1182,15 +1165,9 @@ TEST_F(ParameterizedFp8GemmRewriteTest, ; CHECK-DAG: "operand_precision":["DEFAULT","DEFAULT"] ; CHECK-DAG: } )"; - if (IsRocm() && GetToolkitVersion() < se::SemanticVersion{6, 2, 0}) { - checks += - R"(; CHECK-GCN-DAG: "epilogue":"DEFAULT" -)"; - } else { - checks += - R"(; CHECK-DAG: "epilogue":"BIAS_GELU" + checks += + R"(; CHECK-DAG: "epilogue":"BIAS_GELU" )"; - } checks += R"(; CHECK: } )"; @@ -1257,15 +1234,9 @@ TEST_F(ParameterizedFp8GemmRewriteTest, ; CHECK-NEXT: [[P3:%[^ ]+]] = bf16[] parameter(3) ; CHECK-NEXT: [[XS1:%[^ ]+]] = f32[] convert([[P3]]) )"; - if (IsRocm() && GetToolkitVersion() < se::SemanticVersion{6, 2, 0}) { - checks += - R"(; CHECK-GCN-NEXT: [[OUT:%[^ ]+]] = (f32[16,16]{1,0}, s8[{{[0-9]+}}]{0}) custom-call([[P0]], [[P1_TRANSPOSE]], [[XS]], [[XS1]]), -)"; - } else { - checks += - R"(; CHECK-NEXT: [[OUT:%[^ ]+]] = (bf16[16,16]{1,0}, s8[{{[0-9]+}}]{0}) custom-call([[P0]], [[P1_TRANSPOSE]], [[XS]], [[XS1]]), + checks += + R"(; CHECK-NEXT: [[OUT:%[^ ]+]] = (bf16[16,16]{1,0}, s8[{{[0-9]+}}]{0}) custom-call([[P0]], [[P1_TRANSPOSE]], [[XS]], [[XS1]]), )"; - } checks += R"(; CHECK: custom_call_target="__cublas$lt$matmul$f8", ; CHECK: backend_config={ ; CHECK-DAG: "alpha_real":1 @@ -1281,13 +1252,8 @@ TEST_F(ParameterizedFp8GemmRewriteTest, ; CHECK-DAG: "operand_precision":["DEFAULT","DEFAULT"] ; CHECK-DAG: } )"; - if (IsRocm() && GetToolkitVersion() < se::SemanticVersion{6, 2, 0}) { - checks += R"(; CHECK-GCN-DAG: "epilogue":"DEFAULT" -)"; - } else { - checks += R"(; CHECK-DAG: "epilogue":"GELU" + checks += R"(; CHECK-DAG: "epilogue":"GELU" )"; - } checks += R"(; CHECK: } )"; RunAndFilecheckHloRewrite( From 17bf4fb9506641b8166dca77d86272bbf62463d9 Mon Sep 17 00:00:00 2001 From: Levon Ter-Grigoryan Date: Wed, 30 Sep 2026 01:52:41 -0700 Subject: [PATCH 45/58] [XLA:GPU] Add REDUCE_SCATTER to xla_gpu_experimental_use_collective_kernels. PiperOrigin-RevId: 990832583 --- .../gpu/transforms/collectives/collective_ops_utils.cc | 3 +++ third_party/xla/xla/debug_options_flags.cc | 3 ++- third_party/xla/xla/xla.proto | 3 ++- 3 files changed, 7 insertions(+), 2 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/transforms/collectives/collective_ops_utils.cc b/third_party/xla/xla/backends/gpu/transforms/collectives/collective_ops_utils.cc index babfecbc8046f2..e7908f2f561da9 100644 --- a/third_party/xla/xla/backends/gpu/transforms/collectives/collective_ops_utils.cc +++ b/third_party/xla/xla/backends/gpu/transforms/collectives/collective_ops_utils.cc @@ -451,6 +451,9 @@ absl::StatusOr> OpcodesForTritonCollectives( case xla::DebugOptions::COLLECTIVE_KERNEL_ALL_GATHER: instructions_to_annotate.insert(HloOpcode::kAllGather); break; + case xla::DebugOptions::COLLECTIVE_KERNEL_REDUCE_SCATTER: + instructions_to_annotate.insert(HloOpcode::kReduceScatter); + break; default: return absl::InvalidArgumentError(absl::StrFormat( "Unsupported collective: %s", diff --git a/third_party/xla/xla/debug_options_flags.cc b/third_party/xla/xla/debug_options_flags.cc index 4ff61a22dd4f24..3ebbccccb08f31 100644 --- a/third_party/xla/xla/debug_options_flags.cc +++ b/third_party/xla/xla/debug_options_flags.cc @@ -3331,7 +3331,8 @@ void MakeDebugOptionsFlags(std::vector* flag_list, debug_options->xla_gpu_experimental_use_collective_kernels()), "Experimental: comma-separated filter of collective ops that should use " "custom kernels (e.g. Triton one-shot / two-shot) instead of NCCL. " - "Accepted values: ALL_REDUCE, ALL_GATHER (case-insensitive; the " + "Accepted values: ALL_REDUCE, ALL_GATHER, REDUCE_SCATTER " + "(case-insensitive; the " "COLLECTIVE_KERNEL_ prefix may be omitted). Supports +/- " "incremental modifiers (e.g. +ALL_REDUCE,-ALL_GATHER). The deprecated " "--xla_gpu_unsupported_use_all_reduce_one_shot_kernel flag also adds " diff --git a/third_party/xla/xla/xla.proto b/third_party/xla/xla/xla.proto index 4a56ebe77e5e97..68433174538819 100644 --- a/third_party/xla/xla/xla.proto +++ b/third_party/xla/xla/xla.proto @@ -107,6 +107,7 @@ message DebugOptions { COLLECTIVE_KERNEL_INVALID = 0; COLLECTIVE_KERNEL_ALL_REDUCE = 1; COLLECTIVE_KERNEL_ALL_GATHER = 2; + COLLECTIVE_KERNEL_REDUCE_SCATTER = 3; } enum LibNvJitLinkMode { @@ -1129,7 +1130,7 @@ message DebugOptions { // Experimental: filter specifying which collective operations should use // custom kernels (e.g. Triton one-shot / two-shot) instead of NCCL. - // Accepted values: "all-reduce", "all-gather". + // Accepted values: "all-reduce", "all-gather", "reduce-scatter". // For legacy support, the deprecated // --xla_gpu_unsupported_use_all_reduce_one_shot_kernel flag also adds // all-reduce to this filter. From 0d920ad9698eed033a45bb769ad31aa57710b8d1 Mon Sep 17 00:00:00 2001 From: Akhil Goel Date: Wed, 30 Sep 2026 01:57:24 -0700 Subject: [PATCH 46/58] PR #49300: [XLA:GPU] Inline scan computation to_apply calls Imported from GitHub PR https://github.com/openxla/xla/pull/49300 Scan rewrite lowering converts scan operations into call-based expressions. However, the GPU kernel emitter cannot compute indexing maps for Call operations, causing compilation failures for GPU kernels containing scan. Exact error: ``` E0000 00:00:1789415358.746826 1768463 indexing_analysis.cc:1748] ComputeOutputToInputIndexing is not implemented for opcode call F0000 00:00:1789415358.746906 1768463 computation_partitioner.cc:269] Check failed: operand_maps.size() == 1 (0 vs. 1) ``` This PR inlines the scan computation calls, resolving the indexing analysis failure. CUDA and ROCm backends avoid this issue because they run CallInliner in their device-specific Convolution Canonicalization optimization pipelines that execute before the kernel code generation. This issue came to light after https://github.com/openxla/xla/commit/33d1a586292b8567da377210e71948909358d792 stopped converting scan operations to custom calls. Copybara import of the project: -- ac1221a33121f1a07207d0aad698d85f43ae3bf0 by Akhil Goel : Inline scan calls -- 76a47249786fc94a1b18854444e6b5ba94e72f7b by Akhil Goel : Add CallInliner pass Merging this change closes #49300 PiperOrigin-RevId: 990834717 --- third_party/xla/xla/service/gpu/gpu_compiler.cc | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/third_party/xla/xla/service/gpu/gpu_compiler.cc b/third_party/xla/xla/service/gpu/gpu_compiler.cc index 21b539b21b6e10..52e162c50f1d67 100644 --- a/third_party/xla/xla/service/gpu/gpu_compiler.cc +++ b/third_party/xla/xla/service/gpu/gpu_compiler.cc @@ -924,6 +924,11 @@ absl::Status RunOptimizationPasses( } pipeline.AddPass(); + // AssociativeScanRewriter generates call instructions for scan bodies and the + // emitter cannot compute indexing maps for the Call opcode. Inline the calls + // so the emitter can compute indexing maps. + pipeline.AddPass(); + DynamicPadderOptions dynamic_padder_options; switch (debug_options.xla_gpu_shape_checks()) { From e1d6d57ed5fa7c8e9762595e6668db9beafd06eb Mon Sep 17 00:00:00 2001 From: Mikhail Goncharov Date: Wed, 30 Sep 2026 02:07:04 -0700 Subject: [PATCH 47/58] [XLA:GPU] tighter constraint checks in indexing map 1. when constructed from a map of interval constraints we have not checked if any of the constraints is unsatisfiable - now we sent it trough ctor that calls AddConstraint and handles that for us. 2. added a check if const constraint is satisfiable, e.g. we can get something like `4 in [0, 3]` that should make the whole construct invalid. PiperOrigin-RevId: 990839235 --- third_party/xla/xla/hlo/analysis/BUILD | 1 + .../xla/xla/hlo/analysis/indexing_map.cc | 22 ++------ .../xla/xla/hlo/analysis/indexing_map.h | 13 ++--- .../analysis/indexing_map_serialization.cc | 6 +- .../xla/xla/hlo/analysis/indexing_map_test.cc | 56 +++++++++++++++++-- 5 files changed, 69 insertions(+), 29 deletions(-) diff --git a/third_party/xla/xla/hlo/analysis/BUILD b/third_party/xla/xla/hlo/analysis/BUILD index 952d475c49271a..921167e5c66c23 100644 --- a/third_party/xla/xla/hlo/analysis/BUILD +++ b/third_party/xla/xla/hlo/analysis/BUILD @@ -681,6 +681,7 @@ xla_cc_test( "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/types:span", "@com_google_googletest//:gtest", + "@llvm-project//llvm:Support", "@llvm-project//mlir:IR", ], ) diff --git a/third_party/xla/xla/hlo/analysis/indexing_map.cc b/third_party/xla/xla/hlo/analysis/indexing_map.cc index 856499b94cf052..846b7a3158507f 100644 --- a/third_party/xla/xla/hlo/analysis/indexing_map.cc +++ b/third_party/xla/xla/hlo/analysis/indexing_map.cc @@ -899,16 +899,8 @@ IndexingMap::IndexingMap( std::vector range_vars, std::vector rt_vars, const llvm::MapVector& constraints) - : symbolic_map_(symbolic_map), - dim_vars_(std::move(dimensions)), - range_vars_(std::move(range_vars)), - rt_vars_(std::move(rt_vars)), - constraints_(constraints) { - if (!VerifyVariableIntervals() || !VerifyConstraintIntervals()) { - ResetToKnownEmpty(); - return; - } -} + : IndexingMap(symbolic_map, std::move(dimensions), std::move(range_vars), + std::move(rt_vars), constraints.getArrayRef()) {} IndexingMap IndexingMap::FromTensorSizes( SymbolicMap symbolic_map, absl::Span dim_upper_bounds, @@ -1276,6 +1268,10 @@ bool SymbolicExprSimplifier::SimplifyConstraintExprs(IndexingMap& map) { // Skip constraints that are always satisfied. Interval evaluated_range = range_evaluator_->ComputeExpressionRange(simplified); + if (!evaluated_range.Intersect(range).IsFeasible()) { + map.ResetToKnownEmpty(); + return true; + } if (evaluated_range.upper <= range.upper && evaluated_range.lower >= range.lower) { to_remove.push_back(expr); @@ -1597,12 +1593,6 @@ bool IndexingMap::VerifyVariableIntervals() { }); } -bool IndexingMap::VerifyConstraintIntervals() { - return llvm::all_of(constraints_, [](const auto& constraint) { - return constraint.second.IsFeasible(); - }); -} - SmallBitVector IndexingMap::RemoveUnusedVars() { if (IsUndefined()) { return {}; diff --git a/third_party/xla/xla/hlo/analysis/indexing_map.h b/third_party/xla/xla/hlo/analysis/indexing_map.h index 3682f4061e17d2..b16129507338ee 100644 --- a/third_party/xla/xla/hlo/analysis/indexing_map.h +++ b/third_party/xla/xla/hlo/analysis/indexing_map.h @@ -241,6 +241,11 @@ class IndexingMap { // satisfies both constraints. bool IsKnownEmpty() const { return is_known_empty_; } + // Resets the indexing map to the canonical "known" empty indexing map, i.e. + // (d0...)[s0...]{r0...} -> (0...) symbolic map. + // Does not change the number of symbols, dimensions or results. + void ResetToKnownEmpty(); + bool IsUndefined() const { return symbolic_map_ == SymbolicMap(); } // Removes unused symbols from the `symbolic_map_` and constraints. @@ -289,17 +294,9 @@ class IndexingMap { // Returns true if simplification was performed. bool MergeModConstraints(); - // Resets the indexing map to the canonical "known" empty indexing map, i.e. - // (d0...)[s0...]{r0...} -> (0...) symbolic map. - // Does not change the number of symbols, dimensions or results. - void ResetToKnownEmpty(); - // Verify if all intervals for DimVars, RangeVars and RTVars are feasible. bool VerifyVariableIntervals(); - // Verify if all intervals for constraints. - bool VerifyConstraintIntervals(); - SymbolicMap symbolic_map_; // A dimension variable represents a dimension of a tensor or a GPU grid. diff --git a/third_party/xla/xla/hlo/analysis/indexing_map_serialization.cc b/third_party/xla/xla/hlo/analysis/indexing_map_serialization.cc index b18881ea3b1657..da392d93996698 100644 --- a/third_party/xla/xla/hlo/analysis/indexing_map_serialization.cc +++ b/third_party/xla/xla/hlo/analysis/indexing_map_serialization.cc @@ -590,7 +590,11 @@ std::string ToString(const SymbolicMap& symbolic_map, absl::Span range_names, absl::Span rt_names) { CHECK_EQ(dim_names.size(), symbolic_map.GetNumDims()); - CHECK_EQ(range_names.size() + rt_names.size(), symbolic_map.GetNumSymbols()); + CHECK_EQ(range_names.size() + rt_names.size(), symbolic_map.GetNumSymbols()) + << absl::StrCat("range_names size (", range_names.size(), + ") + rt_names size (", rt_names.size(), + ") != num_symbols in symbolic map (", + symbolic_map.GetNumSymbols(), ")"); std::string s; llvm::raw_string_ostream ss(s); diff --git a/third_party/xla/xla/hlo/analysis/indexing_map_test.cc b/third_party/xla/xla/hlo/analysis/indexing_map_test.cc index f9c917215e4bdd..d18eb7ce8ff95c 100644 --- a/third_party/xla/xla/hlo/analysis/indexing_map_test.cc +++ b/third_party/xla/xla/hlo/analysis/indexing_map_test.cc @@ -27,6 +27,7 @@ limitations under the License. #include "absl/hash/hash_testing.h" #include "absl/strings/string_view.h" #include "absl/types/span.h" +#include "llvm/ADT/MapVector.h" #include "mlir/IR/MLIRContext.h" #include "xla/hlo/analysis/indexing_map_serialization.h" #include "xla/hlo/analysis/indexing_test_utils.h" @@ -401,6 +402,44 @@ TEST_F(IndexingMapTest, EXPECT_THAT(indexing_map, MatchIndexingMap("KNOWN EMPTY")); } +TEST_F(IndexingMapTest, MapVectorConstructorUnsatisfiableConstraints) { + llvm::MapVector constraints; + // Add unsatisfiable constraint. + constraints.insert( + {CreateSymbolicConstant(-1, &mlir_context_), Interval{0, 1}}); + IndexingMap indexing_map(ParseSymbolicMap("(d0) -> (d0)", &mlir_context_), + /*dimensions=*/{IndexingMap::Variable{0, 0}}, + /*range_vars=*/{}, /*rt_vars=*/{}, constraints); + EXPECT_THAT(indexing_map, MatchIndexingMap("KNOWN EMPTY")); +} + +TEST_F(IndexingMapTest, MapVectorConstructorOneOfConstraintsIsUnsatisfiable) { + llvm::MapVector constraints; + constraints.insert( + {CreateSymbolicConstant(-1, &mlir_context_), Interval{0, 1}}); + constraints.insert({CreateDimExpr(0, &mlir_context_), Interval{0, 5}}); + IndexingMap indexing_map(ParseSymbolicMap("(d0) -> (d0)", &mlir_context_), + /*dimensions=*/{IndexingMap::Variable{0, 10}}, + /*range_vars=*/{}, /*rt_vars=*/{}, constraints); + EXPECT_THAT(indexing_map, MatchIndexingMap("KNOWN EMPTY")); + EXPECT_TRUE(indexing_map.GetSymbolicConstraints().empty()); +} + +TEST_F(IndexingMapTest, MapVectorConstructorConstraintAffectsSimplification) { + llvm::MapVector constraints; + constraints.insert({CreateDimExpr(0, &mlir_context_), Interval{0, 7}}); + IndexingMap indexing_map( + ParseSymbolicMap("(d0) -> (d0 mod 16)", &mlir_context_), + /*dimensions=*/{IndexingMap::Variable{0, 31}}, + /*range_vars=*/{}, /*rt_vars=*/{}, constraints); + EXPECT_TRUE(indexing_map.Simplify()); + EXPECT_THAT(indexing_map, MatchIndexingMap(R"( + (d0) -> (d0), + domain: + d0 in [0, 7] + )")); +} + TEST_F(IndexingMapTest, RemoveUnusedVars_ConstraintUsesDim) { // This constraint cannot be removed, because it contains a dimension. auto indexing_map = Parse(R"( @@ -616,19 +655,28 @@ TEST_F(IndexingMapTest, ConstraintIntervalSimplification_Sum) { (d0) -> (d0), domain: d0 in [0, 99], - d0 mod 8 + 5 in [50, 54] + d0 mod 8 + 5 in [6, 10] )"); EXPECT_TRUE(indexing_map.Simplify()); - // TODO: b/459357586 - This should be infeasible, since d0 mod 8 should be in - // [0, 7]. EXPECT_THAT(ToString(indexing_map), MatchIndexingString(R"( (d0) -> (d0), domain: d0 in [0, 99], - d0 mod 8 in [45, 49] + d0 mod 8 in [1, 5] )")); } +TEST_F(IndexingMapTest, ConstraintIntervalSimplification_SumInfeasible) { + auto indexing_map = Parse(R"( + (d0) -> (d0), + domain: + d0 in [0, 99], + d0 mod 8 + 5 in [50, 54] + )"); + EXPECT_TRUE(indexing_map.Simplify()); + EXPECT_THAT(indexing_map, MatchIndexingMap("KNOWN EMPTY")); +} + TEST_F(IndexingMapTest, Simplifier_Mod1) { auto indexing_map = Parse(R"( (d0) -> (d0), From 25acadd55c435d91683d0fda7d366031b1b21fd8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Eusebio=20Dur=C3=A1n=20Monta=C3=B1a?= Date: Wed, 30 Sep 2026 02:19:31 -0700 Subject: [PATCH 48/58] [NFC] Compute GPU instruction annotation titles and metadata once per instruction. `InstructionAnnotation` and `GetInstructionAnnotationMetadata` where being computed twice for the same values. PiperOrigin-RevId: 990844885 --- .../xla/backends/gpu/runtime/annotation.cc | 30 ++++++++----------- 1 file changed, 12 insertions(+), 18 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/runtime/annotation.cc b/third_party/xla/xla/backends/gpu/runtime/annotation.cc index af9ff4e116e475..544bc25c02b8f7 100644 --- a/third_party/xla/xla/backends/gpu/runtime/annotation.cc +++ b/third_party/xla/xla/backends/gpu/runtime/annotation.cc @@ -548,12 +548,10 @@ static std::string MakeInstructionTitle(absl::string_view prefix, return title; } -static std::string MakeInstructionDetails(const HloInstruction& inst) { - // Collect instruction metadata as a key-value suffix that can be parsed by - // XProf. - InstructionAnnotationMetadata metadata = - GetInstructionAnnotationMetadata(inst); - +// Formats instruction metadata as a key-value suffix that can be parsed by +// XProf. +static std::string MakeInstructionDetails( + const InstructionAnnotationMetadata& metadata) { std::string details; auto append = [&](absl::string_view key, std::string value) { if (!value.empty()) { @@ -583,17 +581,12 @@ static std::string MakeInstructionDetails(const HloInstruction& inst) { return details; } -static std::string MakeInstructionName(absl::string_view prefix, - const HloInstruction& inst, - TraceAnnotationLevel annotation_level) { - std::string name = MakeInstructionTitle(prefix, inst); - if (annotation_level < TraceAnnotationLevel::kDetailed) { - return name; +static std::string MakeInstructionName( + absl::string_view title, const InstructionAnnotationMetadata& metadata) { + if (!title.empty() && title.back() == '#') { + title.remove_suffix(1); } - - name.pop_back(); - absl::StrAppend(&name, MakeInstructionDetails(inst), "#"); - return name; + return absl::StrCat(title, MakeInstructionDetails(metadata), "#"); } InstructionAnnotation::InstructionAnnotation( @@ -601,8 +594,6 @@ InstructionAnnotation::InstructionAnnotation( TraceAnnotationLevel annotation_level) : nvtx_name_str_(MakeInstructionTitle( module_annotation.longest_op_name_prefix(), inst)), - xprof_name_str_(MakeInstructionName( - module_annotation.longest_op_name_prefix(), inst, annotation_level)), nvtx_name_(RegisterString(nvtx_name_str_)) { payload_ = Basic{ RegisterString(InstructionAsString(inst)), @@ -611,11 +602,13 @@ InstructionAnnotation::InstructionAnnotation( RegisterString("\n" + CalledInstructionsAsString(inst)), }; if (annotation_level < TraceAnnotationLevel::kDetailed) { + xprof_name_str_ = nvtx_name_str_; return; } InstructionAnnotationMetadata metadata = GetInstructionAnnotationMetadata(inst); + xprof_name_str_ = MakeInstructionName(nvtx_name_str_, metadata); payload_ = Detailed{ std::move(std::get(payload_)), @@ -783,6 +776,7 @@ ModuleAnnotations::ModuleAnnotations(const HloModule& mod, // Loop through `mod` and populate `instructions` with the information we // want to attach to individual instruction ranges. + instructions.reserve(mod.instruction_count()); for (const HloComputation* computation : mod.computations()) { for (const HloInstruction* inst : computation->instructions()) { // e.g. inst.name is "fusion.6", inst.opcode is "kFusion" and called From 9fe191d3e5eab719bf757d2de403ff761fd882bc Mon Sep 17 00:00:00 2001 From: Levon Ter-Grigoryan Date: Wed, 30 Sep 2026 03:09:57 -0700 Subject: [PATCH 49/58] [XLA:GPU] Support disabled VMM API case. This is needed to support fractional vGPUs (see https://github.com/openxla/xla/issues/49252). Collective fusions are disabled in this case and all collectives should fallback on CPU initiated NCCL without RMA access. PiperOrigin-RevId: 990867700 --- .../backends/gpu/tests/all_reduce_e2e_test.cc | 83 +++++++++++++++++++ third_party/xla/xla/debug_options_flags.cc | 7 ++ .../xla/xla/service/gpu/gpu_compiler.cc | 4 +- .../cuda/cuda_device_allocator.cc | 21 ++++- .../cuda/cuda_device_allocator.h | 4 + .../cuda/cuda_device_allocator_test.cc | 35 ++++++++ .../xla/stream_executor/cuda/cuda_executor.cc | 44 ++++++---- third_party/xla/xla/xla.proto | 6 +- 8 files changed, 183 insertions(+), 21 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/tests/all_reduce_e2e_test.cc b/third_party/xla/xla/backends/gpu/tests/all_reduce_e2e_test.cc index 97780efd450c51..72f210749a19ce 100644 --- a/third_party/xla/xla/backends/gpu/tests/all_reduce_e2e_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/all_reduce_e2e_test.cc @@ -1383,5 +1383,88 @@ TEST_F(AllReduceCollectiveKernelTest, } } +TEST_F(AllReduceCollectiveKernelTest, + TritonOneShotAllReduceFallsBackToNcclWhenVmmDisabled) { + constexpr int64_t kNumReplicas = 2; + if (!CheckDeviceCount(kNumReplicas)) { + return; + } + + constexpr absl::string_view kHloText = R"( + HloModule module, replica_count=2 + + add { + lhs = f32[] parameter(0) + rhs = f32[] parameter(1) + ROOT add = f32[] add(lhs, rhs) + } + + ENTRY entry { + param = f32[1024] parameter(0) + ROOT result = f32[1024] all-reduce(param), to_apply=add, replica_groups={{0,1}} + } + )"; + + Literal input_r0 = + LiteralUtil::CreateR1(std::vector(1024, 1.0f)); + Literal input_r1 = + LiteralUtil::CreateR1(std::vector(1024, 2.0f)); + std::vector> args = {{&input_r0}, {&input_r1}}; + Literal expected = + LiteralUtil::CreateR1(std::vector(1024, 3.0f)); + + for (bool enable_command_buffer : {false, true}) { + ASSERT_OK_AND_ASSIGN( + tsl::testing::TemporaryDirectory dump_dir, + tsl::testing::TemporaryDirectory::CreateForCurrentTestcase()); + + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kHloText, kNumReplicas)); + DebugOptions& debug_options = + module->mutable_config().mutable_debug_options(); + debug_options.set_xla_gpu_experimental_vmm_disabled(true); + debug_options.set_xla_gpu_all_reduce_combine_threshold_bytes(0); + debug_options.set_xla_dump_to(dump_dir.path()); + if (!enable_command_buffer) { + debug_options.clear_xla_gpu_enable_command_buffer(); + } else { + debug_options.set_xla_gpu_graph_min_graph_size(1); + } + + ASSERT_OK_AND_ASSIGN(ExecutionResult result, + ExecuteReplicated(std::move(module), args)); + + // Verify that the HLO instruction was annotated as a Triton one-shot + // collective kernel (`KERNEL_STRATEGY_TRITON_ONE_SHOT`), while skipping + // collective fusion so that it lowers to NCCL (`kAllReduceStart`). + VerifyOneShotAllReduce(result.optimized_module); + + ASSERT_OK_AND_ASSIGN( + CommandBufferThunkCounts one_shot, + CountThunksInDump(dump_dir.path(), "kCollectiveKernel")); + EXPECT_EQ(one_shot.in_command_buffer + one_shot.outside_command_buffer, 0); + + ASSERT_OK_AND_ASSIGN(CommandBufferThunkCounts nccl, + CountThunksInDump(dump_dir.path(), "kAllReduce")); + EXPECT_EQ(nccl.in_command_buffer + nccl.outside_command_buffer, 1); + + ASSERT_EQ(result.results.size(), kNumReplicas); + for (int i = 0; i < kNumReplicas; ++i) { + EXPECT_TRUE(LiteralTestUtil::Equal(expected, result.results[i])) + << "Mismatch at replica " << i + << " (enable_command_buffer=" << enable_command_buffer << ")"; + } + + ASSERT_OK_AND_ASSIGN(std::vector second_results, + ExecuteReplicated(result.executable.get(), args)); + ASSERT_EQ(second_results.size(), kNumReplicas); + for (int i = 0; i < kNumReplicas; ++i) { + EXPECT_TRUE(LiteralTestUtil::Equal(expected, second_results[i])) + << "Mismatch at replica " << i << " on second execution" + << " (enable_command_buffer=" << enable_command_buffer << ")"; + } + } +} + } // namespace } // namespace xla diff --git a/third_party/xla/xla/debug_options_flags.cc b/third_party/xla/xla/debug_options_flags.cc index 3ebbccccb08f31..d93b0f97f76b85 100644 --- a/third_party/xla/xla/debug_options_flags.cc +++ b/third_party/xla/xla/debug_options_flags.cc @@ -361,6 +361,7 @@ DebugOptions DefaultDebugOptionsIgnoringFlags() { opts.set_xla_gpu_enable_nccl_user_buffers_in_default_space(false); opts.set_xla_gpu_enable_allocator_spatial_partitioning(true); opts.set_xla_gpu_experimental_enable_nccl_symmetric_buffers(false); + opts.set_xla_gpu_experimental_vmm_disabled(false); opts.set_xla_gpu_experimental_emit_collective_reduce(false); opts.set_xla_gpu_enable_nccl_comm_splitting(true); opts.set_xla_gpu_nccl_init_max_rank_per_root_ratio(0); @@ -3733,6 +3734,12 @@ void MakeDebugOptionsFlags(std::vector* flag_list, &DebugOptions::set_xla_gpu_experimental_scaled_dot_with_triton), debug_options->xla_gpu_experimental_scaled_dot_with_triton(), "If true, use the Triton emitter for scaled dot.")); + flag_list->push_back(tsl::Flag( + "xla_gpu_experimental_vmm_disabled", + bool_setter_for(&DebugOptions::set_xla_gpu_experimental_vmm_disabled), + debug_options->xla_gpu_experimental_vmm_disabled(), + "If true, disables CUDA Virtual Memory Management (VMM) APIs for device " + "memory allocation and collective fusion.")); flag_list->push_back(tsl::Flag( "xla_cpu_collective_call_warn_stuck_timeout_seconds", diff --git a/third_party/xla/xla/service/gpu/gpu_compiler.cc b/third_party/xla/xla/service/gpu/gpu_compiler.cc index 52e162c50f1d67..012b05457db289 100644 --- a/third_party/xla/xla/service/gpu/gpu_compiler.cc +++ b/third_party/xla/xla/service/gpu/gpu_compiler.cc @@ -1544,7 +1544,9 @@ void AddCollectiveCombinerPasses( // so that SolLatencyEstimator and the thunk emitter can consume it. pipeline.AddPass( gpu_topology, /*is_multimem_enabled=*/false); - pipeline.AddPass(gpu_topology); + if (!opts.xla_gpu_experimental_vmm_disabled()) { + pipeline.AddPass(gpu_topology); + } } } diff --git a/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator.cc b/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator.cc index 9b102f027b2c3b..63ea39b673c2ba 100644 --- a/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator.cc +++ b/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator.cc @@ -206,13 +206,22 @@ static CUmemAccessDesc GetAccessDesc(int device) { return descriptor; } -// Allocates device memory using CUDA Virtual Memory Management (VMM) APIs. +// Allocates device memory using CUDA Virtual Memory Management (VMM) APIs, +// or falls back to cuMemAlloc when VMM is disabled. // Returns (virtual_address, padded_size, allocation_handle). static absl::StatusOr> AllocateDeviceMemory(StreamExecutor* executor, const CudaDeviceAllocator::Options& options, uint64_t size) { std::unique_ptr activation = executor->Activate(); + if (!options.use_vmm) { + CUdeviceptr result = 0; + ABSL_RETURN_IF_ERROR(cuda::ToStatus(cuMemAlloc(&result, size))); + void* ptr = absl::bit_cast(result); + XLA_VLOG_DEVICE(3, executor->device_ordinal()) + << "Allocated legacy ptr=" << ptr << " size: " << size; + return std::make_tuple(ptr, size, /*handle=*/0); + } CUdevice device; ABSL_RETURN_IF_ERROR( @@ -390,6 +399,16 @@ void DeallocateDeviceMemory(StreamExecutor* executor, void* ptr, CUmemGenericAllocationHandle handle) { XLA_VLOG_DEVICE(3, executor->device_ordinal()) << "Deallocating " << ptr << " padded size: " << padded_size; + if (handle == 0) { + std::unique_ptr activation = executor->Activate(); + CUdeviceptr pointer = absl::bit_cast(ptr); + absl::Status status = cuda::ToStatus(cuMemFree(pointer)); + if (!status.ok()) { + XLA_LOG_DEVICE(ERROR, executor->device_ordinal()) + << "Failed to free device memory at " << ptr << ": " << status; + } + return; + } ExecutorVmmState* state = GetExecutorVmmState(executor); { diff --git a/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator.h b/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator.h index b68839341df834..17912a529185a8 100644 --- a/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator.h +++ b/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator.h @@ -53,6 +53,10 @@ class CudaDeviceAllocator : public MemoryAllocator { // Whether to mark allocations as GPUDirect RDMA capable. bool enable_rdma = false; + + // Whether to use CUDA Virtual Memory Management (VMM) APIs. If false, + // falls back to legacy cuMemAlloc / cuMemFree APIs. + bool use_vmm = true; }; explicit CudaDeviceAllocator(StreamExecutor* executor); diff --git a/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator_test.cc b/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator_test.cc index f444f26e1f86a6..43bdd6a39d9891 100644 --- a/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator_test.cc +++ b/third_party/xla/xla/stream_executor/cuda/cuda_device_allocator_test.cc @@ -161,5 +161,40 @@ INSTANTIATE_TEST_SUITE_P(RdmaSupport, CudaDeviceAllocatorTest, return info.param ? "RdmaEnabled" : "RdmaDisabled"; }); +TEST(CudaDeviceAllocatorNonVmmTest, AllocateMemcpyAndFreeWithoutVmm) { + ASSERT_OK_AND_ASSIGN(Platform * platform, + PlatformManager::PlatformWithName("CUDA")); + ASSERT_OK_AND_ASSIGN(StreamExecutor * executor, + platform->ExecutorForDevice(0)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr stream, + executor->CreateStream()); + + CudaDeviceAllocator::Options options; + options.use_vmm = false; + CudaDeviceAllocator allocator(executor, options); + + constexpr int kSize = 1024; + ASSERT_OK_AND_ASSIGN(std::unique_ptr allocation, + allocator.Allocate(kSize)); + ASSERT_NE(allocation, nullptr); + EXPECT_NE(allocation->address().opaque(), nullptr); + EXPECT_EQ(allocation->address().size(), kSize); + + std::vector host_src(kSize); + for (int i = 0; i < kSize; i++) { + host_src[i] = static_cast(i); + } + + DeviceAddress addr( + DeviceAddressBase(allocation->address().opaque(), kSize)); + ASSERT_OK(stream->MemcpyH2D(absl::MakeConstSpan(host_src), &addr)); + + std::vector host_dst(kSize, 0); + ASSERT_OK(stream->MemcpyD2H(addr, absl::MakeSpan(host_dst))); + ASSERT_OK(stream->BlockHostUntilDone()); + + EXPECT_EQ(host_src, host_dst); +} + } // namespace } // namespace stream_executor::gpu diff --git a/third_party/xla/xla/stream_executor/cuda/cuda_executor.cc b/third_party/xla/xla/stream_executor/cuda/cuda_executor.cc index 0c78ca77902207..87be1ade40207f 100644 --- a/third_party/xla/xla/stream_executor/cuda/cuda_executor.cc +++ b/third_party/xla/xla/stream_executor/cuda/cuda_executor.cc @@ -975,13 +975,17 @@ CudaExecutor::CreateMemoryAllocator(MemorySpace type) { absl::Status CudaExecutor::Init() { ABSL_ASSIGN_OR_RETURN(device_, GetDevice(device_ordinal())); + const bool vmm_disabled = + xla::GetDebugOptionsFromFlags().xla_gpu_experimental_vmm_disabled(); - ABSL_ASSIGN_OR_RETURN(bool is_vmm_supported, IsVmmSupported(device_)); - if (!is_vmm_supported) { - return absl::InternalError(absl::StrFormat( - "Device %d does not support CUDA Virtual Memory Management (VMM). " - "VMM is required for device memory allocation in XLA.", - device_ordinal())); + if (!vmm_disabled) { + ABSL_ASSIGN_OR_RETURN(bool is_vmm_supported, IsVmmSupported(device_)); + if (!is_vmm_supported) { + return absl::InternalError(absl::StrFormat( + "Device %d does not support CUDA Virtual Memory Management (VMM). " + "VMM is required for device memory allocation in XLA.", + device_ordinal())); + } } ABSL_ASSIGN_OR_RETURN(is_multicast_supported_, IsMulticastSupported(device_)); @@ -1005,18 +1009,22 @@ absl::Status CudaExecutor::Init() { peer_access_cache_[i] = CanEnablePeerAccess(device_, i); } - ABSL_ASSIGN_OR_RETURN(device_allocator_options_, - QueryDeviceAllocatorOptions(device_)); - device_allocator_options_.enable_peer_access = absl::c_any_of( - peer_access_cache_, [](const auto& p) { return p.second; }); - - // Disable fabric handle if there are no active P2P NVLinks — using - // FABRIC+POSIX_FD without a cluster causes allocation failures. - if (device_allocator_options_.enable_fabric_handle && - !GetDeviceDescription().device_interconnect_info().is_in_cluster()) { - XLA_VLOG_DEVICE(2, device_ordinal()) - << "Disable fabric handle on non-cluster machine."; - device_allocator_options_.enable_fabric_handle = false; + if (vmm_disabled) { + device_allocator_options_.use_vmm = false; + } else { + ABSL_ASSIGN_OR_RETURN(device_allocator_options_, + QueryDeviceAllocatorOptions(device_)); + device_allocator_options_.enable_peer_access = absl::c_any_of( + peer_access_cache_, [](const auto& p) { return p.second; }); + + // Disable fabric handle if there are no active P2P NVLinks — using + // FABRIC+POSIX_FD without a cluster causes allocation failures. + if (device_allocator_options_.enable_fabric_handle && + !GetDeviceDescription().device_interconnect_info().is_in_cluster()) { + XLA_VLOG_DEVICE(2, device_ordinal()) + << "Disable fabric handle on non-cluster machine."; + device_allocator_options_.enable_fabric_handle = false; + } } device_allocator_ = diff --git a/third_party/xla/xla/xla.proto b/third_party/xla/xla/xla.proto index 68433174538819..be9d911eff5dd8 100644 --- a/third_party/xla/xla/xla.proto +++ b/third_party/xla/xla/xla.proto @@ -1143,6 +1143,10 @@ message DebugOptions { optional bool xla_gpu_experimental_use_ragged_dot_grouped_gemm = 501; + // If true, disables CUDA Virtual Memory Management (VMM) APIs for device + // memory allocation and collective fusion. + optional bool xla_gpu_experimental_vmm_disabled = 551; + // If true, PTX compilation will fail if a kernel spills registers. // This is meant for debugging and only applies to CUDA PTX compilation. optional bool xla_gpu_fail_ptx_compilation_on_register_spilling = 353; @@ -1884,7 +1888,7 @@ message DebugOptions { // Enables Raft for stable TopK. optional bool xla_gpu_experimental_enable_raft_for_stable_topk = 546; - // Next id: 551 + // Next id: 552 // Extra options to pass to the compilation backend (e.g. LLVM); specific // interpretation of these values is left to the backend. From 645e3ca34317c099a30d7e59d5700f2fbc3bfe28 Mon Sep 17 00:00:00 2001 From: Sohaib Iftikhar Date: Wed, 30 Sep 2026 03:29:25 -0700 Subject: [PATCH 50/58] [XLA:GPU]: Move memcpy inside all-gather Replace host-launched D2D copy to scratch with in-kernel copies to scratch. Each block copies T/R of a tile to the scratch where T is the size of a tile and R is the number of ranks (world_size). Asymmetric block barriers, ensure that consumer blocks wait for all producer blocks to finish before emitting the output copy. PiperOrigin-RevId: 990876953 --- .../gpu/codegen/triton/collective_emitter.cc | 298 +++++++++++++----- .../codegen/triton/collective_emitter_test.cc | 66 ++++ .../triton/tests/collectives/all_gather.hlo | 76 +++-- .../fusion_emitter_shared_dialect_test.cc | 9 +- .../tests/stable_hlo_to_triton_lowering.mlir | 105 ++++++ .../xla/xla/backends/gpu/runtime/BUILD | 2 + .../xla/backends/gpu/runtime/all_gather.cc | 9 +- .../xla/xla/backends/gpu/runtime/all_gather.h | 15 +- .../gpu/runtime/all_gather_build_info_test.cc | 39 +++ third_party/xla/xla/backends/gpu/tests/BUILD | 3 + .../backends/gpu/tests/all_gather_e2e_test.cc | 156 +++++++++ .../xla/xla/codegen/xtile/codegen/BUILD | 2 - .../codegen/xtile/codegen/emitter_helpers.cc | 62 +--- 13 files changed, 662 insertions(+), 180 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/collective_emitter.cc b/third_party/xla/xla/backends/gpu/codegen/triton/collective_emitter.cc index b6cdb9217cf98a..47a6059a6cc719 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/collective_emitter.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/collective_emitter.cc @@ -376,15 +376,22 @@ mlir::LogicalResult PopulateReductionComputation( } // Emits code that reads the last barrier signal value posted by this block from -// its own slot in the local rank's signal buffer (`SignalBuffers[rank] -// [block_id * world_size + rank]`) and returns `last + 1` as the signal value -// for this launch's first barrier. It replaces the host-provided invocation -// count argument (see `CollectiveCodegenConfig::device_sync_count`). +// its own slot in the local rank's signal buffer and returns `last + 1` as the +// signal value for this launch's first barrier. It replaces the host-provided +// invocation count argument (see `CollectiveCodegenConfig::device_sync_count`). // -// `BlockBarrierOp` writes the barrier signal value into -// `SignalBuffers[0..world_size][block_id * world_size + rank]` (after a local -// CTA barrier in `AllReduce`, or from warp 0 in `AllGather`), and no other rank -// or block ever writes to `SignalBuffers[rank][block_id * world_size + rank]`. +// Arguments: +// - `signal_buffers`: `!tt.ptr` table of the per-rank signal buffers. +// - `rank`: i32 rank of this device in the replica group. +// - `block_id`: i32 program id of this block. +// - `world_size`: number of ranks in the replica group. +// - `signal_slot`, `signal_stride`: See `BlockBarrierOp` for details. +// If not set, we use the default indexing of +// SignalBuffers[rank][block_id * world_size + rank]. +// +// After a local CTA barrier, `BlockBarrierOp` writes the barrier signal value +// into this slot of `SignalBuffers[0..world_size]`, and no other rank or block +// ever writes to this slot of `SignalBuffers[rank]`. // For two-shot `AllReduce`, the second barrier writes `signal_value + 1`, so // the slot advances by 2 across launches and adding 1 at entry yields the next // launch's first signal value. Signal buffers are allocated and zeroed once per @@ -393,7 +400,11 @@ mlir::LogicalResult PopulateReductionComputation( mlir::Value EmitDeviceInvocationCount(mlir::ImplicitLocOpBuilder& b, mlir::Value signal_buffers, mlir::Value rank, mlir::Value block_id, - int64_t world_size) { + int64_t world_size, + mlir::Value signal_slot = nullptr, + int64_t signal_stride = 0) { + CHECK_EQ(signal_slot != nullptr, signal_stride > 0) + << "signal_slot and signal_stride must be set together"; const auto ptr_to_i32_type = ttir::PointerType::get(b.getI32Type(), kGlobalAddressSpace); const auto ptr_to_i64_type = @@ -407,14 +418,30 @@ mlir::Value EmitDeviceInvocationCount(mlir::ImplicitLocOpBuilder& b, /*isVolatile=*/false); mlir::Value local_signal_buffer = ttir::IntToPtrOp::create(b, ptr_to_i32_type, local_signal_buffer_i64); - // SignalBuffers[rank][block_id * world_size + rank] - mlir::Value counter_index = arith::AddIOp::create( - b, - arith::MulIOp::create( - b, block_id, - arith::ConstantOp::create( - b, b.getI32IntegerAttr(static_cast(world_size)))), - rank); + mlir::Value world_size_op = arith::ConstantOp::create( + b, b.getI32IntegerAttr(static_cast(world_size))); + mlir::Value counter_index; + if (signal_slot != nullptr) { + // Computes: + // counter_index = ((block_id - signal_slot * signal_stride + + // rank * signal_stride) * world_size) + signal_slot + // See LowerBlockBarrierOp in triton_xla_lower_block_barrier_pass.cc for + // details on the asymmetric barrier indexing. + mlir::Value stride_op = arith::ConstantOp::create( + b, b.getI32IntegerAttr(static_cast(signal_stride))); + mlir::Value slot_offset = arith::MulIOp::create(b, signal_slot, stride_op); + mlir::Value base_block = arith::SubIOp::create(b, block_id, slot_offset); + mlir::Value rank_offset = arith::MulIOp::create(b, rank, stride_op); + mlir::Value target_block = + arith::AddIOp::create(b, base_block, rank_offset); + mlir::Value target_block_offset = + arith::MulIOp::create(b, target_block, world_size_op); + counter_index = arith::AddIOp::create(b, target_block_offset, signal_slot); + } else { + // SignalBuffers[rank][block_id * world_size + rank] + counter_index = arith::AddIOp::create( + b, arith::MulIOp::create(b, block_id, world_size_op), rank); + } mlir::Value counter_ptr = ttir::AddPtrOp::create( b, ptr_to_i32_type, local_signal_buffer, counter_index); // Volatile to make sure that every launch reads the counter from memory. @@ -486,17 +513,19 @@ mlir::Value EmitRemoteBufferPtr(mlir::ImplicitLocOpBuilder& b, // ranks. All threads in the block should have completed their writes before we // proceed with a block barrier. // Otherwise, remote ranks might start reading the data before it is ready. -void EmitBlockBarrier(mlir::ImplicitLocOpBuilder& b, mlir::Value signal_buffers, - mlir::Value rank, mlir::Value signal_value, - int64_t world_size, mlir::Value signal_slot = nullptr, - int64_t signal_stride = 0) { +void EmitBlockBarrier( + mlir::ImplicitLocOpBuilder& b, mlir::Value signal_buffers, mlir::Value rank, + mlir::Value signal_value, int64_t world_size, + mlir::Value signal_slot = nullptr, int64_t signal_stride = 0, + mtx::BarrierMode barrier_mode = mtx::BarrierMode::kSymmetric) { mlir::triton::gpu::BarrierOp::create(b, mlir::triton::gpu::AddrSpace::Local); mtx::BlockBarrierOp::create( b, signal_buffers, rank, signal_value, signal_slot, b.getI32IntegerAttr(static_cast(world_size)), signal_stride > 0 ? b.getI32IntegerAttr(static_cast(signal_stride)) - : mlir::IntegerAttr()); + : mlir::IntegerAttr(), + mtx::BarrierModeAttr::get(b.getContext(), barrier_mode)); } class AllReduceEmitter { @@ -1186,27 +1215,9 @@ absl::StatusOr> GetCollectiveUnmanagedKernelArguments( const HloInstruction* root = computation->root_instruction(); switch (root->opcode()) { case HloOpcode::kAllReduce: + case HloOpcode::kAllGather: return GetRemoteBufferUnmanagedKernelArguments( computation, Cast(root)); - case HloOpcode::kAllGather: { - // AllGather only needs the 3 metadata args (rank, signal_value, - // signal_buffers). No per-parameter scratch buffer args because the - // input parameter already maps to the symmetric scratch buffer via - // the pointer table mechanism in RequiredReplicaIdBounds. - const int32_t num_devices = Cast(root) - ->device_list() - ->num_devices_per_group(); - std::vector unmanaged_arguments; - unmanaged_arguments.reserve(kNumCollectiveMetadataArgs); - // rank and signal_value - unmanaged_arguments.push_back(ShapeUtil::MakeShape(S32, {})); - unmanaged_arguments.push_back(ShapeUtil::MakeShape(S32, {})); - // signal_buffers: pointer-to-pointer table - static constexpr int32_t kMaxBlocksPerGrid = 32; - unmanaged_arguments.push_back( - ShapeUtil::MakeShape(S32, {num_devices, kMaxBlocksPerGrid})); - return unmanaged_arguments; - } default: return std::vector(); } @@ -1223,15 +1234,6 @@ absl::StatusOr AddCollectiveMetadataArguments( fn_arg_types.push_back( ttir::PointerType::get(b.getI64Type(), kGlobalAddressSpace)); - // For AllGather, the input parameter already maps to the symmetric scratch - // buffer via RequiredReplicaIdBounds/SelectBufferOp, so we don't add - // per-parameter scratch buffer opaque args. Only AllReduce (and future ops - // that need explicit remote buffer pointers) add them. - const HloInstruction* root = hlo_computation->root_instruction(); - if (root->opcode() == HloOpcode::kAllGather) { - return kNumCollectiveMetadataArgs; - } - for (HloInstruction* p : hlo_computation->parameter_instructions()) { (void)p; // Also add the remote/scratch buffers for collectives. @@ -1283,12 +1285,20 @@ namespace { struct AllGatherRewriteContext { mlir::stablehlo::AllGatherOp op; - xtile::EntryFuncOp entry_func; xtile::ExtractTileOp input_extract; - xtile::SelectBufferOp pull_select; + // Local input buffer read by `input_extract`. + mlir::TypedValue input_buffer; mlir::Value rank_arg; mlir::Value signal_buffers_arg; + // !tt.ptr table of the gathered parameter's per-rank remote buffers. + mlir::Value remote_buffers_arg; int64_t world_size = 0; + int64_t gather_dim = 0; + int64_t split_dim = -1; + int64_t signal_stride = 0; + PrimitiveType element_type = PrimitiveType::PRIMITIVE_TYPE_INVALID; + llvm::ArrayRef per_rank_shape; + llvm::ArrayRef tile_sizes; }; absl::StatusOr BuildAllGatherContext( @@ -1302,12 +1312,21 @@ absl::StatusOr BuildAllGatherContext( return absl::InvalidArgumentError( "AllGather op must be in an XTile entry function."); } + // The peer rank is derived from the program id (see EmitPeerReplicaId). This + // requires one tile per program, i.e. no tile loop around the op. + if (op->getParentOp() != entry_func.getOperation()) { + return absl::InvalidArgumentError( + "AllGather op must be directly nested in the XTile entry function."); + } + // Same layout as AllReduce: the collective metadata arguments followed by one + // remote buffer argument per fusion parameter. auto num_opaque_attr = entry_func->getAttrOfType("num_opaque_args"); if (!num_opaque_attr || - num_opaque_attr.getInt() < kNumCollectiveMetadataArgs) { + num_opaque_attr.getInt() <= kNumCollectiveMetadataArgs) { return absl::InvalidArgumentError( - "EntryFuncOp does not have collective metadata arguments."); + "EntryFuncOp must have the collective metadata arguments followed by " + "the remote buffer arguments."); } mlir::Value input_tile = op.getOperand(0); @@ -1325,12 +1344,23 @@ absl::StatusOr BuildAllGatherContext( "AllGather operand must be defined by xtile::ExtractTileOp."); } - auto pull_select = llvm::dyn_cast_if_present( - input_extract.getSource().getDefiningOp()); - if (!pull_select) { + // The gathered operand must be read from a fusion parameter. Its argument + // number selects the matching remote buffer argument. + auto input_arg = + mlir::dyn_cast(input_extract.getSource()); + const int64_t num_parameters = + num_opaque_attr.getInt() - kNumCollectiveMetadataArgs; + if (!input_arg || + input_arg.getOwner()->getParentOp() != entry_func.getOperation() || + input_arg.getArgNumber() >= num_parameters) { + return absl::InvalidArgumentError( + "AllGather operand must be extracted from a fusion parameter."); + } + // `input_extract` is redirected to the peer's buffer below, so the local + // tile must not have other users. + if (!input_extract->hasOneUse() || !input_tile.hasOneUse()) { return absl::InvalidArgumentError( - "AllGather ExtractTileOp source must be defined by " - "xtile::SelectBufferOp."); + "AllGather input tile must only be used by the AllGather op."); } ABSL_ASSIGN_OR_RETURN(auto replica_groups, @@ -1340,23 +1370,136 @@ absl::StatusOr BuildAllGatherContext( "AllGather replica groups must not be empty."); } + PrimitiveType element_type = xla::ConvertMlirTypeToPrimitiveType( + input_extract.getType().getElementType()); + if (element_type == PrimitiveType::PRIMITIVE_TYPE_INVALID) { + return absl::InvalidArgumentError( + "Could not convert AllGather element type to PrimitiveType."); + } + AllGatherRewriteContext ctx; ctx.op = op; - ctx.entry_func = entry_func; ctx.input_extract = input_extract; - ctx.pull_select = pull_select; + ctx.input_buffer = input_extract.getSource(); ctx.world_size = replica_groups->num_devices_per_group(); + ctx.gather_dim = op.getAllGatherDim(); + ctx.element_type = element_type; + ctx.per_rank_shape = ctx.input_buffer.getType().getShape(); + ctx.tile_sizes = input_extract.getFullTileShape(); + + // Every tile must be gathered from a single rank (see EmitPeerReplicaId). + const int64_t per_rank_size = ctx.per_rank_shape[ctx.gather_dim]; + const int64_t tile_size = ctx.tile_sizes[ctx.gather_dim]; + if (per_rank_size % tile_size != 0) { + return absl::InvalidArgumentError(absl::StrFormat( + "AllGather tile size (%d) must divide the per-rank size (%d) along the " + "gather dimension.", + tile_size, per_rank_size)); + } int32_t total_args = entry_func.getNumArguments(); int32_t metadata_args_start = total_args - num_opaque_attr.getInt() - kNumTileIndexArgs; // Layout: opaque[0]=rank, opaque[1]=invocation count (unused, the signal - // value comes from a counter in device memory), opaque[2]=signal_buffers. + // value comes from a counter in device memory), opaque[2]=signal_buffers, + // opaque[3 + i]=remote buffers of parameter i (same as AllReduce). ctx.rank_arg = entry_func.getArgument(metadata_args_start); ctx.signal_buffers_arg = entry_func.getArgument(metadata_args_start + 2); + ctx.remote_buffers_arg = + entry_func.getArgument(metadata_args_start + kNumCollectiveMetadataArgs + + input_arg.getArgNumber()); + + // Find a dimension to split the tile size across `world_size` to balance + // the in-kernel D2D subtile copy work across participating ranks. + // We prefer splitting along `gather_dim`, but fallback to the innermost + // cleanly divisible dimension. If none are found, split_dim remains -1 + // (unsplit). + if (ctx.tile_sizes[ctx.gather_dim] >= ctx.world_size && + ctx.tile_sizes[ctx.gather_dim] % ctx.world_size == 0) { + ctx.split_dim = ctx.gather_dim; + } else { + for (int64_t d = static_cast(ctx.tile_sizes.size()) - 1; d >= 0; + --d) { + if (ctx.tile_sizes[d] >= ctx.world_size && + ctx.tile_sizes[d] % ctx.world_size == 0) { + ctx.split_dim = d; + break; + } + } + } + + ctx.signal_stride = + ctx.per_rank_shape[ctx.gather_dim] / ctx.tile_sizes[ctx.gather_dim]; + for (int64_t j = ctx.gather_dim + 1; + j < static_cast(ctx.per_rank_shape.size()); ++j) { + ctx.signal_stride *= + xla::CeilOfRatio(ctx.per_rank_shape[j], ctx.tile_sizes[j]); + } return ctx; } +// Returns the rank whose output slice contains this program's tile. Programs +// enumerate output tiles in row-major order (same assumption as +// `signal_stride` and the asymmetric barrier), so each run of `signal_stride` +// consecutive programs covers one rank's slice along `gather_dim`. +mlir::Value EmitPeerReplicaId(mlir::ImplicitLocOpBuilder& b, + const AllGatherRewriteContext& ctx, + mlir::Value block_id) { + mlir::Value stride_op = arith::ConstantOp::create( + b, b.getI32IntegerAttr(static_cast(ctx.signal_stride))); + mlir::Value world_size_op = arith::ConstantOp::create( + b, b.getI32IntegerAttr(static_cast(ctx.world_size))); + mlir::Value tile_group = arith::DivUIOp::create(b, block_id, stride_op); + return arith::RemUIOp::create(b, tile_group, world_size_op); +} + +// Returns the active slot of `rank`'s remote buffer as a memref with the type +// of the local input buffer. +mlir::Value EmitRemoteBufferMemref(mlir::ImplicitLocOpBuilder& b, + const AllGatherRewriteContext& ctx, + mlir::Value rank, + mlir::Value buffer_offset) { + const mlir::MemRefType memref_type = ctx.input_buffer.getType(); + mlir::Value ptr = + EmitRemoteBufferPtr(b, ctx.remote_buffers_arg, rank, buffer_offset, + memref_type.getElementType()); + return mtx::PtrToMemrefOp::create(b, memref_type, ptr); +} + +void CopySubtileToScratch(mlir::ImplicitLocOpBuilder& b, + AllGatherRewriteContext& ctx, + mlir::Value peer_replica_id, + mlir::Value buffer_offset) { + llvm::SmallVector stage_offsets( + ctx.input_extract.getOffsets().begin(), + ctx.input_extract.getOffsets().end()); + llvm::SmallVector subtile_shape(ctx.tile_sizes.begin(), + ctx.tile_sizes.end()); + if (ctx.split_dim != -1) { + const int64_t subtile_size = ctx.tile_sizes[ctx.split_dim] / ctx.world_size; + mlir::Value subtile_size_val = + arith::ConstantIndexOp::create(b, subtile_size); + mlir::Value peer_replica_idx = + arith::IndexCastOp::create(b, b.getIndexType(), peer_replica_id); + mlir::Value subtile_off = + arith::MulIOp::create(b, peer_replica_idx, subtile_size_val); + stage_offsets[ctx.split_dim] = + arith::AddIOp::create(b, stage_offsets[ctx.split_dim], subtile_off); + subtile_shape[ctx.split_dim] = subtile_size; + } + + auto subtile_tensor_type = mlir::RankedTensorType::get( + subtile_shape, ctx.input_extract.getType().getElementType()); + auto subtile = xtile::ExtractTileOp::create( + b, subtile_tensor_type, ctx.input_buffer, stage_offsets, subtile_shape, + ctx.input_extract.getStrides()); + + mlir::Value own_scratch_slot = + EmitRemoteBufferMemref(b, ctx, ctx.rank_arg, buffer_offset); + xtile::InsertTileOp::create(b, subtile, own_scratch_slot, stage_offsets, + subtile_shape, ctx.input_extract.getStrides()); +} + } // namespace mlir::LogicalResult RewriteAllGather(mlir::stablehlo::AllGatherOp op, @@ -1367,20 +1510,31 @@ mlir::LogicalResult RewriteAllGather(mlir::stablehlo::AllGatherOp op, } AllGatherRewriteContext& ctx = *maybe_ctx; - mlir::ImplicitLocOpBuilder builder(ctx.pull_select.getLoc(), rewriter); - builder.setInsertionPoint(ctx.pull_select); + mlir::ImplicitLocOpBuilder builder(op.getLoc(), rewriter); + builder.setInsertionPoint(ctx.input_extract); mlir::Value block_id = ttir::GetProgramIdOp::create(builder, 0); + mlir::Value peer_replica_id = EmitPeerReplicaId(builder, ctx, block_id); mlir::Value signal_value = EmitDeviceInvocationCount( - builder, ctx.signal_buffers_arg, ctx.rank_arg, block_id, ctx.world_size); - // Inter-block barrier via signal flags. This blocks until all - // remote ranks have also signaled. - mtx::BlockBarrierOp::create(builder, ctx.signal_buffers_arg, ctx.rank_arg, - signal_value, /*signal_slot=*/nullptr, - builder.getI32IntegerAttr(ctx.world_size), - /*signal_stride=*/nullptr); - + builder, ctx.signal_buffers_arg, ctx.rank_arg, block_id, ctx.world_size, + /*signal_slot=*/peer_replica_id, ctx.signal_stride); + mlir::Value buffer_offset = EmitDoubleBufferOffset( + builder, signal_value, Product(ctx.per_rank_shape), ctx.element_type); + + // Copy this program's subtile of the local tile to this rank's remote buffer. + CopySubtileToScratch(builder, ctx, peer_replica_id, buffer_offset); + // Wait until the peer has staged the whole tile. + EmitBlockBarrier(builder, ctx.signal_buffers_arg, ctx.rank_arg, signal_value, + ctx.world_size, /*signal_slot=*/peer_replica_id, + ctx.signal_stride, mtx::BarrierMode::kConsumerSymmetric); + // Pull the tile from the peer's remote buffer. + mlir::Value pull_scratch_slot = + EmitRemoteBufferMemref(builder, ctx, peer_replica_id, buffer_offset); + rewriter.modifyOpInPlace(ctx.input_extract, [&]() { + ctx.input_extract.getSourceMutable().assign(pull_scratch_slot); + }); rewriter.replaceOp(op, op.getOperand(0)); + return mlir::success(); } diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/collective_emitter_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/collective_emitter_test.cc index 7b555f05f9d8ea..c966f2f87b0d2f 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/collective_emitter_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/collective_emitter_test.cc @@ -496,6 +496,72 @@ TEST_F(CollectiveBlockLevelConfigTest, AllGatherBlockLevelConfig) { EXPECT_THAT(block_level_config.output_tiles(0).sizes(), ElementsAre(2048)); } +TEST_F(CollectiveEmitterTest, AllGatherGetCollectiveUnmanagedKernelArguments) { + constexpr absl::string_view kAllGatherHloStr = R"( + HloModule test + ENTRY test_computation { + param_0 = f32[32768] parameter(0) + ROOT all-gather = f32[65536] all-gather(param_0), replica_groups={{0,1}}, + dimensions={0} + } + )"; + ASSERT_OK_AND_ASSIGN( + ModuleWithFusion module_with_fusion, + BuildModuleWithFusion(kAllGatherHloStr, HloOpcode::kAllGather)); + ASSERT_OK_AND_ASSIGN( + const auto unmanaged_arguments, + GetCollectiveUnmanagedKernelArguments(module_with_fusion.FusionInstr())); + ASSERT_EQ(unmanaged_arguments.size(), 4); + // [0]: rank (S32[]) + EXPECT_EQ(unmanaged_arguments[0].dimensions().size(), 0); + // [1]: signal_value (S32[]) + EXPECT_EQ(unmanaged_arguments[1].dimensions().size(), 0); + // [2]: signal_buffers (S32[num_devices, kMaxBlocksPerGrid]) + ASSERT_EQ(unmanaged_arguments[2].dimensions().size(), 2); + EXPECT_EQ(unmanaged_arguments[2].dimensions()[0], 2); + // [3]: remote buffers of param_0 (F32[num_devices, 32768]) + EXPECT_THAT(unmanaged_arguments[3].dimensions(), ElementsAre(2, 32768)); +} + +TEST_F(CollectiveBlockLevelConfigTest, + AllGatherBlockLevelConfigClampsGatherDimToPerRankSize) { + constexpr absl::string_view kAllGatherHloStr = R"( + HloModule test + ENTRY test_computation { + param_0 = f32[1,64] parameter(0) + ROOT all-gather = f32[2,64] all-gather(param_0), replica_groups={{0,1}}, + dimensions={0} + } + )"; + ASSERT_OK_AND_ASSIGN( + ModuleWithFusion module_with_fusion, + BuildModuleWithFusion(kAllGatherHloStr, HloOpcode::kAllGather)); + ASSERT_OK_AND_ASSIGN(const BlockLevelFusionConfig block_level_config, + GetCollectiveBlockLevelFusionConfig( + *gpu_topology_, module_with_fusion.FusionInstr())); + ASSERT_EQ(block_level_config.output_tiles_size(), 1); + EXPECT_THAT(block_level_config.output_tiles(0).sizes(), ElementsAre(1, 64)); +} + +TEST_F(CollectiveBlockLevelConfigTest, AllGatherBlockLevelConfigAtMaxBlocks) { + constexpr absl::string_view kAllGatherHloStr = R"( + HloModule test + ENTRY test_computation { + param_0 = f32[65536] parameter(0) + ROOT all-gather = f32[131072] all-gather(param_0), replica_groups={{0,1}}, + dimensions={0} + } + )"; + ASSERT_OK_AND_ASSIGN( + ModuleWithFusion module_with_fusion, + BuildModuleWithFusion(kAllGatherHloStr, HloOpcode::kAllGather)); + ASSERT_OK_AND_ASSIGN(const BlockLevelFusionConfig block_level_config, + GetCollectiveBlockLevelFusionConfig( + *gpu_topology_, module_with_fusion.FusionInstr())); + // 131072 elements / kAllGatherMaxBlocksPerGrid (64) blocks. + EXPECT_THAT(block_level_config.output_tiles(0).sizes(), ElementsAre(2048)); +} + } // namespace } // namespace xla::gpu diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/tests/collectives/all_gather.hlo b/third_party/xla/xla/backends/gpu/codegen/triton/tests/collectives/all_gather.hlo index b065647800ae85..1b838aa65c5c1a 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/tests/collectives/all_gather.hlo +++ b/third_party/xla/xla/backends/gpu/codegen/triton/tests/collectives/all_gather.hlo @@ -25,37 +25,49 @@ ENTRY entry { } } -// CHECK-DAG: #[[MAP0:.*]] = #xla.indexing_map<"(pid) -> (((pid / 8) mod 8) * 16), domain: pid in [0, 127]"> -// CHECK-DAG: #[[MAP1:.*]] = #xla.indexing_map<"(pid) -> ((pid mod 8) * 16), domain: pid in [0, 127]"> -// CHECK-DAG: #[[MAP2:.*]] = #xla.indexing_map<"(pid) -> (pid / 64), domain: pid in [0, 127]"> -// CHECK-DAG: #[[MAP3:.*]] = #xla.indexing_map<"(pid) -> ((pid / 8) * 16), domain: pid in [0, 127]"> +// f32[128,128] per rank with 16x16 output tiles gives 128 programs. Programs +// 0-63 produce rank 0's slice and 64-127 produce rank 1's slice, so program +// `pid` reads its tile from rank peer = (pid / 64) % 2. -// CHECK: xtile.entry_func @triton_fn( -// CHECK-SAME: %[[ARG0:[a-zA-Z0-9_]+]]: memref<2xi64>, -// CHECK-SAME: %[[ARG1:[a-zA-Z0-9_]+]]: memref<256x128xf32>, -// CHECK-SAME: %[[RANK:[a-zA-Z0-9_]+]]: i32, -// CHECK-SAME: %[[INVOCATION_COUNT:[a-zA-Z0-9_]+]]: i32, -// CHECK-SAME: %[[SIGNAL_BUFFERS:[a-zA-Z0-9_]+]]: !tt.ptr, -// CHECK-SAME: %[[PID:[a-zA-Z0-9_]+]]: index {xla.range = [0 : index, 127 : index]}) +// CHECK-LABEL: xtile.entry_func @triton_fn( +// CHECK-SAME: %[[INPUT:[^:]+]]: memref<128x128xf32>, +// CHECK-SAME: %[[OUTPUT:[^:]+]]: memref<256x128xf32>, +// CHECK-SAME: %[[RANK:[^:]+]]: i32, +// CHECK-SAME: %{{[^:]+}}: i32, +// CHECK-SAME: %[[SIGNAL_BUFFERS:[^:]+]]: !tt.ptr, +// CHECK-SAME: %[[REMOTE_BUFFERS:[^:]+]]: !tt.ptr, +// CHECK-SAME: attributes {num_opaque_args = 4 : i32} +// CHECK-DAG: %[[WORLD_SIZE:.*]] = arith.constant 2 : i32 +// CHECK-DAG: %[[TILES_PER_RANK:.*]] = arith.constant 64 : i32 +// CHECK: %[[PID:.*]] = tt.get_program_id x +// CHECK: %[[SLICE:.*]] = arith.divui %[[PID]], %[[TILES_PER_RANK]] +// CHECK: %[[PEER:.*]] = arith.remui %[[SLICE]], %[[WORLD_SIZE]] -// CHECK-DAG: %[[IDX0:.*]] = xla.apply_indexing #[[MAP0]](%[[PID]]) -// CHECK-DAG: %[[IDX1:.*]] = xla.apply_indexing #[[MAP1]](%[[PID]]) -// CHECK-DAG: %[[IDX2:.*]] = xla.apply_indexing #[[MAP2]](%[[PID]]) -// CHECK-DAG: %[[C1:.*]] = arith.constant 1 : i32 -// CHECK-DAG: %[[WORLD_SIZE:.*]] = arith.constant 2 : i32 -// CHECK: %[[BLOCK_ID:.*]] = tt.get_program_id x -// CHECK: %[[SIGNAL_BUFFER_ADDR:.*]] = tt.addptr %[[SIGNAL_BUFFERS]], %[[RANK]] -// CHECK: %[[SIGNAL_BUFFER_INT:.*]] = tt.load %[[SIGNAL_BUFFER_ADDR]] -// CHECK: %[[SIGNAL_BUFFER:.*]] = tt.int_to_ptr %[[SIGNAL_BUFFER_INT]] -// CHECK: %[[BLOCK_OFFSET:.*]] = arith.muli %[[BLOCK_ID]], %[[WORLD_SIZE]] -// CHECK: %[[COUNTER_IDX:.*]] = arith.addi %[[BLOCK_OFFSET]], %[[RANK]] -// CHECK: %[[COUNTER_PTR:.*]] = tt.addptr %[[SIGNAL_BUFFER]], %[[COUNTER_IDX]] -// CHECK: %[[COUNTER:.*]] = tt.load %[[COUNTER_PTR]] {isVolatile = true} -// CHECK: %[[SIGNAL_VALUE:.*]] = arith.addi %[[COUNTER]], %[[C1]] -// CHECK: triton_xla.block_barrier %[[SIGNAL_BUFFERS]], %[[RANK]], %[[SIGNAL_VALUE]] -// CHECK-SAME: -// CHECK: %[[BUF:.*]] = xtile.select_buffer %[[ARG0]][%[[IDX2]]] : memref<2xi64> -> memref<128x128xf32> -// CHECK: %[[EXT:.*]] = xtile.extract %[[BUF]][%[[IDX0]], %[[IDX1]]] [16, 16] [1, 1] : memref<128x128xf32> -> tensor<16x16xf32> -// CHECK: %[[IDX3:.*]] = xla.apply_indexing #[[MAP3]](%[[PID]]) -// CHECK: xtile.insert %[[EXT]] into %[[ARG1]][%[[IDX3]], %[[IDX1]]] [16, 16] [1, 1] : tensor<16x16xf32> -> memref<256x128xf32> -// CHECK: xtile.return +// The invocation count is loaded from this rank's signal buffer at the slot +// of the asymmetric block barrier: (pid - peer * 64 + rank * 64) * 2 + peer. +// CHECK: %[[SLOT_OFFSET:.*]] = arith.muli %[[PEER]], %[[TILES_PER_RANK]] +// CHECK: %[[BASE_BLOCK:.*]] = arith.subi %[[PID]], %[[SLOT_OFFSET]] +// CHECK: %[[RANK_OFFSET:.*]] = arith.muli %[[RANK]], %[[TILES_PER_RANK]] +// CHECK: %[[TARGET_BLOCK:.*]] = arith.addi %[[BASE_BLOCK]], %[[RANK_OFFSET]] +// CHECK: %[[TARGET_OFFSET:.*]] = arith.muli %[[TARGET_BLOCK]], %[[WORLD_SIZE]] +// CHECK: %[[COUNTER_IDX:.*]] = arith.addi %[[TARGET_OFFSET]], %[[PEER]] +// CHECK: tt.addptr %{{.*}}, %[[COUNTER_IDX]] : !tt.ptr, i32 + +// Stage this program's half of the local tile (8 of 16 rows) in this rank's +// remote buffer. The program 64 apart stages the other half. +// CHECK: %[[HALF:.*]] = xtile.extract %[[INPUT]]{{.*}} [8, 16] [1, 1] +// CHECK: tt.addptr %[[REMOTE_BUFFERS]], %[[RANK]] +// CHECK: %[[OWN_BUFFER:.*]] = triton_xla.ptr_to_memref +// CHECK: xtile.insert %[[HALF]] into %[[OWN_BUFFER]] + +// Wait until the peer has staged the whole tile. +// CHECK: ttg.barrier local +// CHECK: triton_xla.block_barrier %[[SIGNAL_BUFFERS]], %[[RANK]], %{{.*}}, %[[PEER]] +// CHECK-SAME: + +// Copy the tile from the peer's remote buffer to the output. +// CHECK: tt.addptr %[[REMOTE_BUFFERS]], %[[PEER]] +// CHECK: %[[PEER_BUFFER:.*]] = triton_xla.ptr_to_memref +// CHECK: %[[TILE:.*]] = xtile.extract %[[PEER_BUFFER]]{{.*}} [16, 16] [1, 1] +// CHECK: xtile.insert %[[TILE]] into %[[OUTPUT]] +// CHECK: xtile.return diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/tests/fusion_emitter_shared_dialect_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/tests/fusion_emitter_shared_dialect_test.cc index 4b12d7be11a5fe..dd75f9fa03c715 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/tests/fusion_emitter_shared_dialect_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/tests/fusion_emitter_shared_dialect_test.cc @@ -578,12 +578,9 @@ TEST_F(XTileDialectTest, HloAllGatherDotLowering) { EXPECT_OK(CreateXTileIrAndFileCheck(*module->GetComputationWithName("ag_dot"), block_level_parameters, R"( - CHECK: xtile.entry_func @xtile_dialect_fn(%arg0: memref<2xi64> - CHECK: %[[SELECT1:.*]] = xtile.select_buffer %arg0[%{{.*}}] - CHECK-SAME: : memref<2xi64> -> memref<2xi64> - CHECK: %[[SELECT2:.*]] = xtile.select_buffer %[[SELECT1]][%{{.*}}] - CHECK-SAME: : memref<2xi64> -> memref<128x128xf32> - CHECK: %[[LHS_TILE:.*]] = xtile.extract %[[SELECT2]] + CHECK: xtile.entry_func @xtile_dialect_fn(%arg0: memref<128x128xf32>, %arg1: memref<128x128xf32>, %arg2: memref<512x128xf32>, %arg3: index + CHECK-NOT: xtile.select_buffer + CHECK: %[[LHS_TILE:.*]] = xtile.extract %arg0 CHECK: %[[AG1:.*]] = "stablehlo.all_gather"(%[[LHS_TILE]]) CHECK: %[[AG2:.*]] = "stablehlo.all_gather"(%[[AG1]]) CHECK: %[[RHS_TILE:.*]] = xtile.extract %arg1 diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/transforms/tests/stable_hlo_to_triton_lowering.mlir b/third_party/xla/xla/backends/gpu/codegen/triton/transforms/tests/stable_hlo_to_triton_lowering.mlir index 18c7c464421313..fb556e2d59787b 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/transforms/tests/stable_hlo_to_triton_lowering.mlir +++ b/third_party/xla/xla/backends/gpu/codegen/triton/transforms/tests/stable_hlo_to_triton_lowering.mlir @@ -358,6 +358,111 @@ xtile.entry_func @all_reduce_two_shot_3d(%input: memref<1024x512x2xf32>, %output xtile.return } +// CHECK-LABEL: xtile.entry_func @all_gather_one_shot( +xtile.entry_func @all_gather_one_shot(%input: memref<128x128xf32>, %output: memref<256x128xf32>, %device_rank: i32, %signal_value: i32, %signal_buffer: !tt.ptr, %remote_input_buffer: !tt.ptr, %tile_id: index) attributes {num_opaque_args = 4 : i32} { + %c0 = arith.constant 0 : index + %tile = xtile.extract %input[%c0, %c0][16, 16][1, 1] : memref<128x128xf32> -> tensor<16x16xf32> + // CHECK: arith.divui + // CHECK: arith.remui + // CHECK: triton_xla.ptr_to_memref + // CHECK: triton_xla.block_barrier {{.*}} + // CHECK: triton_xla.ptr_to_memref + // CHECK-NOT: stablehlo.all_gather + %all_gather = "stablehlo.all_gather"(%tile) <{all_gather_dim = 0 : i64, replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>}> : (tensor<16x16xf32>) -> tensor<16x16xf32> + xtile.insert %all_gather into %output[%c0, %c0][16, 16][1, 1] : tensor<16x16xf32> -> memref<256x128xf32> + xtile.return +} + +// Neither tile dimension is divisible by world_size, so every program stages +// the whole tile instead of a 1/world_size part of it. +// CHECK-LABEL: xtile.entry_func @all_gather_tile_not_divisible_by_world_size( +// CHECK-SAME: %[[INPUT:[a-zA-Z0-9_]+]]: memref<4x4xf32> +// CHECK-NOT: arith.index_cast +// CHECK: xtile.extract %[[INPUT]]{{.*}} [1, 1] [1, 1] +// CHECK: triton_xla.block_barrier {{.*}} +xtile.entry_func @all_gather_tile_not_divisible_by_world_size(%input: memref<4x4xf32>, %output: memref<8x4xf32>, %device_rank: i32, %signal_value: i32, %signal_buffer: !tt.ptr, %remote_input_buffer: !tt.ptr, %tile_id: index) attributes {num_opaque_args = 4 : i32} { + %c0 = arith.constant 0 : index + %tile = xtile.extract %input[%c0, %c0][1, 1][1, 1] : memref<4x4xf32> -> tensor<1x1xf32> + %all_gather = "stablehlo.all_gather"(%tile) <{all_gather_dim = 0 : i64, replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>}> : (tensor<1x1xf32>) -> tensor<1x1xf32> + xtile.insert %all_gather into %output[%c0, %c0][1, 1][1, 1] : tensor<1x1xf32> -> memref<8x4xf32> + xtile.return +} + +// The gather-dim tile (1) is not divisible by world_size, so the split falls +// back to the innermost divisible dimension: each program stages [1, 8]. +// CHECK-LABEL: xtile.entry_func @all_gather_split_falls_back_to_inner_dim( +// CHECK-SAME: %[[INPUT:[a-zA-Z0-9_]+]]: memref<4x16xf32> +// CHECK: arith.index_cast +// CHECK: xtile.extract %[[INPUT]]{{.*}} [1, 8] [1, 1] +// CHECK: triton_xla.block_barrier {{.*}} +xtile.entry_func @all_gather_split_falls_back_to_inner_dim(%input: memref<4x16xf32>, %output: memref<8x16xf32>, %device_rank: i32, %signal_value: i32, %signal_buffer: !tt.ptr, %remote_input_buffer: !tt.ptr, %tile_id: index) attributes {num_opaque_args = 4 : i32} { + %c0 = arith.constant 0 : index + %tile = xtile.extract %input[%c0, %c0][1, 16][1, 1] : memref<4x16xf32> -> tensor<1x16xf32> + %all_gather = "stablehlo.all_gather"(%tile) <{all_gather_dim = 0 : i64, replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>}> : (tensor<1x16xf32>) -> tensor<1x16xf32> + xtile.insert %all_gather into %output[%c0, %c0][1, 16][1, 1] : tensor<1x16xf32> -> memref<8x16xf32> + xtile.return +} + +// CHECK-LABEL: xtile.entry_func @all_gather_second_parameter( +// CHECK-SAME: %[[INPUT1:[a-zA-Z0-9_]+]]: memref<128x128xf32>, %{{[a-zA-Z0-9_]+}}: memref<256x128xf32> +// CHECK-SAME: %[[REMOTE1:[a-zA-Z0-9_]+]]: !tt.ptr, %{{[a-zA-Z0-9_]+}}: index +xtile.entry_func @all_gather_second_parameter(%input0: memref<128x128xf32>, %input1: memref<128x128xf32>, %output: memref<256x128xf32>, %device_rank: i32, %signal_value: i32, %signal_buffer: !tt.ptr, %remote_input_buffer0: !tt.ptr, %remote_input_buffer1: !tt.ptr, %tile_id: index) attributes {num_opaque_args = 5 : i32} { + %c0 = arith.constant 0 : index + // CHECK: xtile.extract %[[INPUT1]][ + // CHECK: tt.addptr %[[REMOTE1]], %{{.*}} : !tt.ptr, i32 + // CHECK: triton_xla.block_barrier + // CHECK: tt.addptr %[[REMOTE1]], %{{.*}} : !tt.ptr, i32 + %tile = xtile.extract %input1[%c0, %c0][16, 16][1, 1] : memref<128x128xf32> -> tensor<16x16xf32> + %all_gather = "stablehlo.all_gather"(%tile) <{all_gather_dim = 0 : i64, replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>}> : (tensor<16x16xf32>) -> tensor<16x16xf32> + xtile.insert %all_gather into %output[%c0, %c0][16, 16][1, 1] : tensor<16x16xf32> -> memref<256x128xf32> + xtile.return +} + +// CHECK-LABEL: xtile.entry_func @all_gather_in_loop_doesnt_lower( +xtile.entry_func @all_gather_in_loop_doesnt_lower(%input: memref<128x128xf32>, %output: memref<256x128xf32>, %device_rank: i32, %signal_value: i32, %signal_buffer: !tt.ptr, %remote_input_buffer: !tt.ptr, %tile_id: index) attributes {num_opaque_args = 4 : i32} { + %c0 = arith.constant 0 : index + %c1 = arith.constant 1 : index + %c2 = arith.constant 2 : index + scf.for %i = %c0 to %c2 step %c1 { + %tile = xtile.extract %input[%c0, %c0][16, 16][1, 1] : memref<128x128xf32> -> tensor<16x16xf32> + // CHECK: stablehlo.all_gather + %all_gather = "stablehlo.all_gather"(%tile) <{all_gather_dim = 0 : i64, replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>}> : (tensor<16x16xf32>) -> tensor<16x16xf32> + xtile.insert %all_gather into %output[%c0, %c0][16, 16][1, 1] : tensor<16x16xf32> -> memref<256x128xf32> + } + xtile.return +} + +// CHECK-LABEL: xtile.entry_func @all_gather_tile_not_dividing_per_rank_size_doesnt_lower( +xtile.entry_func @all_gather_tile_not_dividing_per_rank_size_doesnt_lower(%input: memref<96x128xf32>, %output: memref<192x128xf32>, %device_rank: i32, %signal_value: i32, %signal_buffer: !tt.ptr, %remote_input_buffer: !tt.ptr, %tile_id: index) attributes {num_opaque_args = 4 : i32} { + %c0 = arith.constant 0 : index + %tile = xtile.extract %input[%c0, %c0][64, 16][1, 1] : memref<96x128xf32> -> tensor<64x16xf32> + // CHECK: stablehlo.all_gather + %all_gather = "stablehlo.all_gather"(%tile) <{all_gather_dim = 0 : i64, replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>}> : (tensor<64x16xf32>) -> tensor<64x16xf32> + xtile.insert %all_gather into %output[%c0, %c0][64, 16][1, 1] : tensor<64x16xf32> -> memref<192x128xf32> + xtile.return +} + +// CHECK-LABEL: xtile.entry_func @all_gather_input_tile_with_other_users_doesnt_lower( +xtile.entry_func @all_gather_input_tile_with_other_users_doesnt_lower(%input: memref<128x128xf32>, %output: memref<256x128xf32>, %device_rank: i32, %signal_value: i32, %signal_buffer: !tt.ptr, %remote_input_buffer: !tt.ptr, %tile_id: index) attributes {num_opaque_args = 4 : i32} { + %c0 = arith.constant 0 : index + %tile = xtile.extract %input[%c0, %c0][16, 16][1, 1] : memref<128x128xf32> -> tensor<16x16xf32> + // CHECK: stablehlo.all_gather + %all_gather = "stablehlo.all_gather"(%tile) <{all_gather_dim = 0 : i64, replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>}> : (tensor<16x16xf32>) -> tensor<16x16xf32> + %sum = arith.addf %all_gather, %tile : tensor<16x16xf32> + xtile.insert %sum into %output[%c0, %c0][16, 16][1, 1] : tensor<16x16xf32> -> memref<256x128xf32> + xtile.return +} + +// CHECK-LABEL: xtile.entry_func @all_gather_without_remote_buffers_arg_doesnt_lower( +xtile.entry_func @all_gather_without_remote_buffers_arg_doesnt_lower(%input: memref<128x128xf32>, %output: memref<256x128xf32>, %device_rank: i32, %signal_value: i32, %signal_buffer: !tt.ptr, %tile_id: index) attributes {num_opaque_args = 3 : i32} { + %c0 = arith.constant 0 : index + %tile = xtile.extract %input[%c0, %c0][16, 16][1, 1] : memref<128x128xf32> -> tensor<16x16xf32> + // CHECK: stablehlo.all_gather + %all_gather = "stablehlo.all_gather"(%tile) <{all_gather_dim = 0 : i64, replica_groups = dense<[[0, 1]]> : tensor<1x2xi64>}> : (tensor<16x16xf32>) -> tensor<16x16xf32> + xtile.insert %all_gather into %output[%c0, %c0][16, 16][1, 1] : tensor<16x16xf32> -> memref<256x128xf32> + xtile.return +} + // CHECK: func @lower_dot_with_warp_specialization_to_triton func.func @lower_dot_with_warp_specialization_to_triton( %arg0: tensor<2x4xf32>, diff --git a/third_party/xla/xla/backends/gpu/runtime/BUILD b/third_party/xla/xla/backends/gpu/runtime/BUILD index 67bd53871d044c..32a2b62df23868 100644 --- a/third_party/xla/xla/backends/gpu/runtime/BUILD +++ b/third_party/xla/xla/backends/gpu/runtime/BUILD @@ -4339,6 +4339,7 @@ xla_cc_test( srcs = ["all_gather_build_info_test.cc"], deps = [ ":all_gather", + ":collective_params", "//xla:shape_util", "//xla:xla_data_proto_cc", "//xla/backends/gpu/target_config", @@ -4346,6 +4347,7 @@ xla_cc_test( "//xla/hlo/testlib:hlo_hardware_independent_test_base", "//xla/service:gpu_topology", "//xla/service/gpu:gpu_device_info_for_tests", + "//xla/service/gpu:launch_dimensions", "//xla/stream_executor:device_description", "//xla/stream_executor:device_description_proto_cc", "//xla/tsl/lib/gtl:int_type", diff --git a/third_party/xla/xla/backends/gpu/runtime/all_gather.cc b/third_party/xla/xla/backends/gpu/runtime/all_gather.cc index f0e781ec3bb726..440b158ba9e0b7 100644 --- a/third_party/xla/xla/backends/gpu/runtime/all_gather.cc +++ b/third_party/xla/xla/backends/gpu/runtime/all_gather.cc @@ -231,19 +231,20 @@ absl::StatusOr CreateAllGatherKernelSpec( CollectiveKernelSpec kernel_spec = { /* .codegen_config= */ { - /* .copy_input_to_scratch= */ true, + /* .copy_input_to_scratch= */ false, /* .input_buffer_specs= */ {{/*requires_multimem=*/false, SymmetricMemoryType::kNone}}, /* .output_buffer_specs= */ {{/*requires_multimem=*/false, SymmetricMemoryType::kNone}}, /* .argument_descriptors= */ - {{KernelArgType::kScratchBuffer, - /*index=*/1}, // scratch buffer as input + {{KernelArgType::kInputBuffer, /*index=*/0}, {KernelArgType::kOutputBuffer, /*index=*/0}, {KernelArgType::kRuntimeRank}, {KernelArgType::kInvocationCount}, {KernelArgType::kScratchBuffer, - /*index=*/0}}, // signal buffers only + /*index=*/0}, // signal buffers + {KernelArgType::kScratchBuffer, + /*index=*/1}}, // remote scratch buffers /* .sync_count_increment= */ 1u}, /* .scratch_buffers= */ {{signal_size, /*requires_multimem=*/false, sym_mem_type, diff --git a/third_party/xla/xla/backends/gpu/runtime/all_gather.h b/third_party/xla/xla/backends/gpu/runtime/all_gather.h index 82e6ebca9486aa..82bc362111c08f 100644 --- a/third_party/xla/xla/backends/gpu/runtime/all_gather.h +++ b/third_party/xla/xla/backends/gpu/runtime/all_gather.h @@ -46,10 +46,10 @@ inline constexpr auto kSupportedAllGatherTypes = std::array{F16, BF16, F32, F64, S8, S16, S32, S64}; // Maximum number of GPU thread-blocks launched per all-gather kernel. -// This constant is shared between the kernel launcher (all_gather.cc) and the -// unmanaged-argument shaper (collective_emitter.cc) so that the signal buffer -// is always sized to match the actual grid. -inline constexpr int64_t kAllGatherMaxBlocksPerGrid = 32; +// This number is set by picking the max number of blocks for all reduce and +// scaling it because the outputs are bigger for all-gather. +// The number is somewhat arbitrary and should be revisited. +inline constexpr int64_t kAllGatherMaxBlocksPerGrid = 64; // Optimal threshold for one-shot all-gather in bytes for the collective kernel. // Base on the experimental results. @@ -97,14 +97,13 @@ LaunchDimensions AllGatherLaunchDimensions( const se::DeviceDescription& device_info); // Creates a CollectiveKernelSpec describing the resource requirements of a -// Triton all-gather kernel. The kernel argument layout is: -// [0] input/scratch buffer pointer table (kScratchBuffer, index 1) +// Triton all-gather kernel. The kernel argument layout matches all-reduce: +// [0] local input buffer (kInputBuffer, index 0) // [1] output buffer (kOutputBuffer, index 0) // [2] runtime rank (kRuntimeRank) // [3] invocation count (kInvocationCount) // [4] signal flags (kScratchBuffer, index 0) -// The runtime performs a D2D copy from the input buffer to the local rank's -// scratch buffer before kernel launch (copy_input_to_scratch=true). +// [5] remote scratch buffer pointer table (kScratchBuffer, index 1) absl::StatusOr CreateAllGatherKernelSpec( const HloInstruction* instr, const LaunchDimensions& launch_dimensions); diff --git a/third_party/xla/xla/backends/gpu/runtime/all_gather_build_info_test.cc b/third_party/xla/xla/backends/gpu/runtime/all_gather_build_info_test.cc index 7e926b4c85d3af..9bd8e2df1ab291 100644 --- a/third_party/xla/xla/backends/gpu/runtime/all_gather_build_info_test.cc +++ b/third_party/xla/xla/backends/gpu/runtime/all_gather_build_info_test.cc @@ -27,6 +27,7 @@ limitations under the License. #include "absl/strings/str_join.h" #include "absl/strings/string_view.h" #include "xla/backends/gpu/runtime/all_gather.h" +#include "xla/backends/gpu/runtime/collective_params.h" #include "xla/backends/gpu/target_config/target_config.h" #include "xla/hlo/ir/hlo_casting_utils.h" #include "xla/hlo/ir/hlo_instruction.h" @@ -36,6 +37,7 @@ limitations under the License. #include "xla/hlo/testlib/hlo_hardware_independent_test_base.h" #include "xla/primitive_util.h" #include "xla/service/gpu/gpu_device_info_for_tests.h" +#include "xla/service/gpu/launch_dimensions.h" #include "xla/service/gpu_topology.h" #include "xla/stream_executor/device_description.h" #include "xla/stream_executor/device_description.pb.h" @@ -222,5 +224,42 @@ TEST_F(BuildAllGatherInfoTest, FailsForLargeInputs) { HasSubstr("only supported for small inputs"))); } +TEST_F(BuildAllGatherInfoTest, + CreateAllGatherKernelSpecMatchesAllReduceArgumentLayout) { + constexpr absl::string_view kModuleStr = R"( + HloModule test + ENTRY test_computation { + param_0 = f32[512] parameter(0) + ROOT all-gather = f32[1024] all-gather(param_0), + dimensions={0}, replica_groups={{0,1}} + } + )"; + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(kModuleStr, 2)); + const HloInstruction* instr = HloHardwareIndependentTestBase::FindInstruction( + module.get(), HloOpcode::kAllGather); + ASSERT_OK_AND_ASSIGN( + CollectiveKernelSpec spec, + CreateAllGatherKernelSpec(instr, LaunchDimensions(4, 128))); + EXPECT_FALSE(spec.codegen_config.copy_input_to_scratch); + ASSERT_EQ(spec.codegen_config.argument_descriptors.size(), 6); + EXPECT_EQ(spec.codegen_config.argument_descriptors[0].type, + KernelArgType::kInputBuffer); + EXPECT_EQ(spec.codegen_config.argument_descriptors[0].index, 0); + EXPECT_EQ(spec.codegen_config.argument_descriptors[1].type, + KernelArgType::kOutputBuffer); + EXPECT_EQ(spec.codegen_config.argument_descriptors[1].index, 0); + EXPECT_EQ(spec.codegen_config.argument_descriptors[2].type, + KernelArgType::kRuntimeRank); + EXPECT_EQ(spec.codegen_config.argument_descriptors[3].type, + KernelArgType::kInvocationCount); + EXPECT_EQ(spec.codegen_config.argument_descriptors[4].type, + KernelArgType::kScratchBuffer); + EXPECT_EQ(spec.codegen_config.argument_descriptors[4].index, 0); + EXPECT_EQ(spec.codegen_config.argument_descriptors[5].type, + KernelArgType::kScratchBuffer); + EXPECT_EQ(spec.codegen_config.argument_descriptors[5].index, 1); +} + } // namespace } // namespace xla::gpu diff --git a/third_party/xla/xla/backends/gpu/tests/BUILD b/third_party/xla/xla/backends/gpu/tests/BUILD index c7fd9b12c2af57..c439f9a6457636 100644 --- a/third_party/xla/xla/backends/gpu/tests/BUILD +++ b/third_party/xla/xla/backends/gpu/tests/BUILD @@ -1708,7 +1708,10 @@ xla_test( ":collective_ops_e2e_test_base", "//xla:literal", "//xla:literal_util", + "//xla:shape_util", + "//xla:xla_data_proto_cc", "//xla:xla_proto_cc", + "//xla/backends/gpu/runtime:all_gather", "//xla/hlo/ir:hlo", "//xla/service/gpu:backend_configs_cc", "//xla/tests:literal_test_util", diff --git a/third_party/xla/xla/backends/gpu/tests/all_gather_e2e_test.cc b/third_party/xla/xla/backends/gpu/tests/all_gather_e2e_test.cc index 60c7b25fb614ec..2b045929ee29a9 100644 --- a/third_party/xla/xla/backends/gpu/tests/all_gather_e2e_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/all_gather_e2e_test.cc @@ -23,6 +23,7 @@ limitations under the License. #include "absl/strings/str_cat.h" #include "absl/strings/str_format.h" #include "absl/strings/string_view.h" +#include "xla/backends/gpu/runtime/all_gather.h" #include "xla/backends/gpu/tests/collective_ops_e2e_test_base.h" #include "xla/hlo/ir/hlo_computation.h" #include "xla/hlo/ir/hlo_instruction.h" @@ -30,11 +31,14 @@ limitations under the License. #include "xla/hlo/ir/hlo_opcode.h" #include "xla/literal.h" #include "xla/literal_util.h" +#include "xla/primitive_util.h" #include "xla/service/gpu/backend_configs.pb.h" +#include "xla/shape_util.h" #include "xla/tests/literal_test_util.h" #include "xla/tsl/platform/test.h" #include "xla/tsl/testing/temporary_directory.h" #include "xla/xla.pb.h" +#include "xla/xla_data.pb.h" namespace xla { namespace { @@ -509,5 +513,157 @@ TEST_F(AllGatherTest, TwentyAllGathersInWhileLoopWithCudaGraphs) { } } +// 2D shape: f32[16, 32] -> f32[16, 64] (gather along dim 1). +TEST_F(AllGatherTest, TwoDimensionalGatherDim1) { + constexpr int32_t kNumReplicas = 2; + if (!CheckDeviceCount(kNumReplicas)) { + return; + } + + constexpr absl::string_view kModuleStr = R"( + HloModule test + ENTRY test_computation { + param_0 = f32[16,32] parameter(0) + ROOT all-gather = f32[16,64] all-gather(param_0), dimensions={1}, + replica_groups={{0,1}} + } + )"; + + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + + Literal input_r0 = LiteralUtil::CreateFull({16, 32}, 10.0f); + Literal input_r1 = LiteralUtil::CreateFull({16, 32}, 20.0f); + + std::vector> args = {{&input_r0}, {&input_r1}}; + ASSERT_OK_AND_ASSIGN(ExecutionResult result, + ExecuteReplicated(std::move(module), args)); + + VerifyOneShotAllGather(result.optimized_module); + + ASSERT_EQ(result.results.size(), kNumReplicas); + + Literal expected = LiteralUtil::CreateFull({16, 64}, 0.0f); + for (int64_t row = 0; row < 16; ++row) { + for (int64_t col = 0; col < 64; ++col) { + expected.Set({row, col}, (col < 32) ? 10.0f : 20.0f); + } + } + + for (int i = 0; i < kNumReplicas; ++i) { + EXPECT_TRUE(LiteralTestUtil::Equal(expected, result.results[i])) + << "Mismatch at replica " << i; + } +} + +// Repeated invocations on the same executable to verify double-buffering +// across alternating slots (signal_value & 1). +TEST_F(AllGatherTest, RepeatedInvocationsDoubleBuffering) { + constexpr int32_t kNumReplicas = 2; + if (!CheckDeviceCount(kNumReplicas)) { + return; + } + + constexpr absl::string_view kModuleStr = R"( + HloModule test + ENTRY test_computation { + param_0 = f32[128] parameter(0) + ROOT all-gather = f32[256] all-gather(param_0), dimensions={0}, + replica_groups={{0,1}} + } + )"; + + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + + Literal input1_r0 = + LiteralUtil::CreateR1(std::vector(128, 1.0f)); + Literal input1_r1 = + LiteralUtil::CreateR1(std::vector(128, 2.0f)); + std::vector> args1 = {{&input1_r0}, {&input1_r1}}; + + ASSERT_OK_AND_ASSIGN(ExecutionResult result1, + ExecuteReplicated(std::move(module), args1)); + VerifyOneShotAllGather(result1.optimized_module); + + for (int iter = 2; iter <= 4; ++iter) { + float v0 = static_cast(iter * 10 + 1); + float v1 = static_cast(iter * 10 + 2); + Literal in_r0 = LiteralUtil::CreateR1(std::vector(128, v0)); + Literal in_r1 = LiteralUtil::CreateR1(std::vector(128, v1)); + std::vector> iter_args = {{&in_r0}, {&in_r1}}; + + ASSERT_OK_AND_ASSIGN( + std::vector iter_results, + ExecuteReplicated(result1.executable.get(), iter_args)); + ASSERT_EQ(iter_results.size(), kNumReplicas); + + std::vector expected_data(256); + for (int i = 0; i < 128; ++i) { + expected_data[i] = v0; + expected_data[i + 128] = v1; + } + Literal expected = LiteralUtil::CreateR1(expected_data); + for (int i = 0; i < kNumReplicas; ++i) { + EXPECT_TRUE(LiteralTestUtil::Equal(expected, iter_results[i])) + << "Mismatch at replica " << i << " on iteration " << iter; + } + } +} + +class AllGatherTypesTest : public AllGatherTest, + public ::testing::WithParamInterface { +}; + +// Same shape as Large2GpuF32, for every element type supported by the +// one-shot kernel. +TEST_P(AllGatherTypesTest, Large2Gpu) { + constexpr int32_t kNumReplicas = 2; + if (!CheckDeviceCount(kNumReplicas)) { + return; + } + + constexpr absl::string_view kModuleStr = R"( + HloModule test + ENTRY test_computation { + param_0 = %1$s[4096] parameter(0) + ROOT all-gather = %1$s[8192] all-gather(param_0), dimensions={0}, + replica_groups={{0,1}} + } + )"; + + const PrimitiveType type = GetParam(); + const std::string module_str = absl::StrFormat( + kModuleStr, primitive_util::LowercasePrimitiveTypeName(type)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(module_str, kNumReplicas)); + + // Rank r contributes the r-th half of the expected output. + ASSERT_OK_AND_ASSIGN(Literal expected, + MakeFakeLiteral(ShapeUtil::MakeShape(type, {8192}))); + Literal input_r0 = expected.Slice({0}, {4096}); + Literal input_r1 = expected.Slice({4096}, {8192}); + + std::vector> args = {{&input_r0}, {&input_r1}}; + ASSERT_OK_AND_ASSIGN(ExecutionResult result, + ExecuteReplicated(std::move(module), args)); + + VerifyOneShotAllGather(result.optimized_module); + + ASSERT_EQ(result.results.size(), kNumReplicas); + for (int i = 0; i < kNumReplicas; ++i) { + EXPECT_TRUE(LiteralTestUtil::Equal(expected, result.results[i])) + << "Mismatch at replica " << i; + } +} + +INSTANTIATE_TEST_SUITE_P( + AllGatherTypes, AllGatherTypesTest, + ::testing::ValuesIn(gpu::kSupportedAllGatherTypes), + [](const ::testing::TestParamInfo& info) { + return std::string( + primitive_util::LowercasePrimitiveTypeName(info.param)); + }); + } // namespace } // namespace xla diff --git a/third_party/xla/xla/codegen/xtile/codegen/BUILD b/third_party/xla/xla/codegen/xtile/codegen/BUILD index cd80e0c8f22bc9..af1f15977030f9 100644 --- a/third_party/xla/xla/codegen/xtile/codegen/BUILD +++ b/third_party/xla/xla/codegen/xtile/codegen/BUILD @@ -68,13 +68,11 @@ cc_library( "@com_google_absl//absl/status:status_macros", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", - "@com_google_absl//absl/strings:str_format", "@com_google_absl//absl/types:span", "@llvm-project//llvm:Support", "@llvm-project//mlir:ArithDialect", "@llvm-project//mlir:ArithUtils", "@llvm-project//mlir:ComplexDialect", - "@llvm-project//mlir:DialectUtils", "@llvm-project//mlir:IR", "@llvm-project//mlir:MathDialect", "@llvm-project//mlir:Support", diff --git a/third_party/xla/xla/codegen/xtile/codegen/emitter_helpers.cc b/third_party/xla/xla/codegen/xtile/codegen/emitter_helpers.cc index c31075340d64b7..e3f3f93a51db99 100644 --- a/third_party/xla/xla/codegen/xtile/codegen/emitter_helpers.cc +++ b/third_party/xla/xla/codegen/xtile/codegen/emitter_helpers.cc @@ -28,7 +28,6 @@ limitations under the License. #include "absl/status/status_macros.h" #include "absl/status/statusor.h" #include "absl/strings/str_cat.h" -#include "absl/strings/str_format.h" #include "absl/strings/str_join.h" #include "absl/strings/string_view.h" #include "absl/types/span.h" @@ -41,7 +40,6 @@ limitations under the License. #include "mlir/Dialect/Arith/Utils/Utils.h" #include "mlir/Dialect/Math/IR/Math.h" #include "mlir/Dialect/Tensor/IR/Tensor.h" -#include "mlir/Dialect/Utils/StaticValueUtils.h" #include "mlir/IR/Builders.h" #include "mlir/IR/BuiltinAttributes.h" #include "mlir/IR/BuiltinOps.h" @@ -936,51 +934,14 @@ Value Bitcast(mlir::ImplicitLocOpBuilder& b, Value value, Type type) { std::move(replica_id_offsets), std::move(replica_id_bounds)); } -absl::StatusOr GetConstantIntValue(mlir::Value value) { - if (std::optional int_value = mlir::getConstantIntValue(value); - int_value.has_value()) { - return int_value.value(); - } - return absl::InternalError(absl::StrFormat( - "Expected constant integer value for replica ID bound, but got: %v", - value)); -} - absl::StatusOr EmitParameterExtract(mlir::ImplicitLocOpBuilder& b, const TileInfo& tile_info, Value arg) { auto tensor_type = mlir::RankedTensorType::get(tile_info.padded_tile_sizes(), tile_info.storage_type()); - mlir::Value source_buffer = arg; - if (!tile_info.replica_id_offsets().empty()) { - const auto& replica_id_offsets = tile_info.replica_id_offsets(); - const auto& replica_id_bounds = tile_info.replica_id_bounds(); - CHECK_EQ(replica_id_offsets.size(), replica_id_bounds.size()); - const int num_replica_dims = replica_id_offsets.size(); - for (int i = 0; i < num_replica_dims - 1; ++i) { - mlir::Value replica_id = replica_id_offsets[i]; - ABSL_ASSIGN_OR_RETURN(int64_t next_bound, - GetConstantIntValue(replica_id_bounds[i + 1])); - mlir::Type next_buffer_type = - mlir::MemRefType::get({next_bound}, b.getI64Type()); - source_buffer = b.create( - next_buffer_type, source_buffer, replica_id); - } - // Final selection to obtain the spatial buffer - mlir::Value replica_id = replica_id_offsets.back(); - ABSL_ASSIGN_OR_RETURN(PrimitiveType element_type, - GetPrimitiveType(tile_info.storage_type())); - xla::Shape spatial_shape = xla::ShapeUtil::MakeShapeWithDenseLayout( - element_type, tile_info.storage_shape(), - tile_info.minor_to_major_layout()); - ABSL_ASSIGN_OR_RETURN(mlir::MemRefType spatial_memref_type, - GetMemRefType(spatial_shape, tile_info.storage_type())); - source_buffer = b.create(spatial_memref_type, - source_buffer, replica_id); - } return xla::xtile::ExtractTileOp::create( - b, tensor_type, source_buffer, tile_info.offsets(), - tile_info.padded_tile_sizes(), tile_info.tile_strides()); + b, tensor_type, arg, tile_info.offsets(), tile_info.padded_tile_sizes(), + tile_info.tile_strides()); } absl::StatusOr EmitScope( @@ -1162,7 +1123,7 @@ absl::StatusOr> GetFnArgTypes( mlir::ImplicitLocOpBuilder& b, const HloFusionInstruction& fusion, absl::Span opaque_args_types, const std::optional& gpu_cc, - const DefaultTileRequirementsVisitor& tile_requirements_visitor) { + const DefaultTileRequirementsVisitor& /*tile_requirements_visitor*/) { SmallVector fn_arg_types; auto hlo_computation = fusion.fused_instructions_computation(); @@ -1170,20 +1131,9 @@ absl::StatusOr> GetFnArgTypes( for (HloInstruction* p : hlo_computation->parameter_instructions()) { ABSL_ASSIGN_OR_RETURN(Type ir_type, GetMlirType(b, p->shape().element_type(), gpu_cc)); - ABSL_ASSIGN_OR_RETURN(SmallVector replica_id_bounds, - tile_requirements_visitor.RequiredReplicaIdBounds(*p)); - if (!replica_id_bounds.empty()) { - // Nested pointer schema for replica dimensions. - // R x S x where R is the number of replica dimensions and S is - // the shape on the local device. In total we have R pointers to - // S-dimensional tensors. - fn_arg_types.push_back( - mlir::MemRefType::get({replica_id_bounds.front()}, b.getI64Type())); - } else { - ABSL_ASSIGN_OR_RETURN(mlir::MemRefType memref_type, - GetMemRefType(p->shape(), ir_type)); - fn_arg_types.push_back(memref_type); - } + ABSL_ASSIGN_OR_RETURN(mlir::MemRefType memref_type, + GetMemRefType(p->shape(), ir_type)); + fn_arg_types.push_back(memref_type); } // Add result types. From 5a1f0fd4ae099a23e269920a23c8f78d25d84142 Mon Sep 17 00:00:00 2001 From: Srishti Srivastava Date: Wed, 30 Sep 2026 03:33:51 -0700 Subject: [PATCH 51/58] [StableHLO] Fix windows link error in evalRunParallel Implemented the fix suggested here: https://github.com/openxla/stablehlo/pull/2975. PiperOrigin-RevId: 990878885 --- .../xla/third_party/stablehlo/temporary.patch | 62 +++++++++++++++++++ 1 file changed, 62 insertions(+) diff --git a/third_party/xla/third_party/stablehlo/temporary.patch b/third_party/xla/third_party/stablehlo/temporary.patch index 678053cd922786..32d1ceba4c11b1 100644 --- a/third_party/xla/third_party/stablehlo/temporary.patch +++ b/third_party/xla/third_party/stablehlo/temporary.patch @@ -2292,6 +2292,68 @@ diff --ruN a/stablehlo/stablehlo/tests/CheckOps.h b/stablehlo/stablehlo/tests/Ch uint64_t min_ulp_difference, uint64_t max_ulp_difference); +diff --ruN a/stablehlo/stablehlo/reference/InterpreterOps.cpp b/stablehlo/stablehlo/reference/InterpreterOps.cpp +--- stablehlo/stablehlo/reference/InterpreterOps.cpp ++++ stablehlo/stablehlo/reference/InterpreterOps.cpp +@@ -22,6 +22,7 @@ + #include + #include + #include ++#include + + #include "llvm/Support/Casting.h" + #include "llvm/Support/Errc.h" +@@ -175,11 +176,14 @@ + SmallVector> programs, SymbolTable& symbolTable, + InterpreterFallback* fallback) { + llvm::DefaultThreadPool threadPool; +- SmallVector>> futures; ++ llvm::ThreadPoolTaskGroup taskGroup(threadPool); + + uint32_t numReplicas = programs.size(); + uint32_t numPartitions = programs[0].size(); + ProcessGrid processGrid(numReplicas, numPartitions, infeed); ++ ++ SmallVector> taskOutputs(numReplicas * ++ numPartitions); + + auto inputsIt = inputs.begin(); + +@@ -187,24 +191,24 @@ + for (uint32_t j = 0; j < numPartitions; ++j) { + auto funcName = programs[i][j]; + auto func = llvm::cast(symbolTable.lookup(funcName)); +- auto evalWrapper = [&](Region& region, ArrayRef args, +- ProcessId processId) { +- Process process{processId, &processGrid}; +- return eval(region, args, fallback, &process, +- /*parent=*/nullptr); +- }; +- + auto numArgs = func.getBody().front().getArguments().size(); + SmallVector args(inputsIt, inputsIt + numArgs); + inputsIt += numArgs; + +- futures.emplace_back(threadPool.async( +- evalWrapper, std::ref(func.getBody()), args, ProcessId{i, j})); ++ uint32_t flatIndex = i * numPartitions + j; ++ taskGroup.async([&, flatIndex, func, args = std::move(args), ++ processId = ProcessId{i, j}]() mutable { ++ Process process{processId, &processGrid}; ++ taskOutputs[flatIndex] = eval(func.getBody(), args, fallback, &process, ++ /*parent=*/nullptr); ++ }); + } + } + ++ taskGroup.wait(); ++ + SmallVector results; +- for (auto& future : futures) results.append(future.get()); ++ for (auto& output : taskOutputs) results.append(output); + // TODO(#1725): Figure out how to test the outfeed queue. + return results; + } diff --ruN a/stablehlo/stablehlo/tests/TestUtils.cpp b/stablehlo/stablehlo/tests/TestUtils.cpp --- stablehlo/stablehlo/tests/TestUtils.cpp +++ stablehlo/stablehlo/tests/TestUtils.cpp From 270a24580ffd7d68e4755af7d7943328a27fe5cd Mon Sep 17 00:00:00 2001 From: Eugene Zhulenev Date: Wed, 30 Sep 2026 03:44:34 -0700 Subject: [PATCH 52/58] PR #49719: [tsl:concurrency] Deduplicate AsyncValue type ids by type name across DSOs Imported from GitHub PR https://github.com/openxla/xla/pull/49719 `AsyncValue` assigns each payload type a 16-bit type id by appending to a process-global `TypeInfoTable` and caching the resulting index in `GetTypeId()`'s function-local static. This is unsafe once XLA is split across dynamically-linked libraries: XLA is built with `-fvisibility=hidden`, which demotes the `GetTypeId()` local static to a per-DSO local. A type registered in more than one DSO therefore appends to the shared table more than once and receives a different id in each DSO, so `IsType()`/`DynCast()` on an `AsyncValue` that crosses a DSO boundary can silently return the wrong answer. Make registration idempotent by type name: `CreateTypeInfoAndReturnTypeIdImpl` now keys a `name->id` map on `typeid(T).name()` and returns the existing id when the same type is registered again, instead of allocating a fresh one. The map lives behind a const-init mutex in the single out-of-line definition of the registration function, so all DSOs that resolve that symbol share one id per type. No caller changes are required. This mirrors MLIR's `TypeID`, whose fallback resolver (r`egisterImplicitTypeID(getTypeName())`) already keys on the type name rather than on registration order, giving a process-stable identity that survives across DSO boundary. Copybara import of the project: -- e94bb1d32f7a06d397a4d33cd4b1a4dd0c07a4e6 by Eugene Zhulenev : [tsl:concurrency] Deduplicate AsyncValue type ids by type name across DSOs Merging this change closes #49719 PiperOrigin-RevId: 990883982 --- third_party/xla/xla/tsl/concurrency/BUILD | 1 + .../xla/xla/tsl/concurrency/async_value.cc | 24 ++++++++++++++++--- .../xla/xla/tsl/concurrency/async_value.h | 13 ++++++++-- 3 files changed, 33 insertions(+), 5 deletions(-) diff --git a/third_party/xla/xla/tsl/concurrency/BUILD b/third_party/xla/xla/tsl/concurrency/BUILD index 05d4ba34f1f95a..ecc6365e405685 100644 --- a/third_party/xla/xla/tsl/concurrency/BUILD +++ b/third_party/xla/xla/tsl/concurrency/BUILD @@ -50,6 +50,7 @@ cc_library( "//xla/tsl/util:safe_reinterpret_cast", "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/base:no_destructor", + "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/container:inlined_vector", "@com_google_absl//absl/functional:any_invocable", "@com_google_absl//absl/status", diff --git a/third_party/xla/xla/tsl/concurrency/async_value.cc b/third_party/xla/xla/tsl/concurrency/async_value.cc index 2ca9faa5890066..c5437bd1eef875 100644 --- a/third_party/xla/xla/tsl/concurrency/async_value.cc +++ b/third_party/xla/xla/tsl/concurrency/async_value.cc @@ -16,30 +16,48 @@ limitations under the License. #include "xla/tsl/concurrency/async_value.h" #include +#include #include #include #include +#include #include +#include "absl/base/attributes.h" +#include "absl/base/const_init.h" #include "absl/base/no_destructor.h" #include "absl/base/optimization.h" +#include "absl/container/flat_hash_map.h" #include "absl/container/inlined_vector.h" #include "absl/functional/any_invocable.h" +#include "absl/strings/string_view.h" +#include "absl/synchronization/mutex.h" #include "absl/synchronization/notification.h" #include "absl/types/span.h" #include "xla/tsl/concurrency/async_value_ref.h" #include "xla/tsl/concurrency/ref_count.h" #include "xla/tsl/platform/logging.h" -#include "tsl/platform/context.h" namespace tsl { uint16_t AsyncValue::CreateTypeInfoAndReturnTypeIdImpl( - const TypeInfo& type_info) { + absl::string_view type_name, const TypeInfo& type_info) { + // Deduplicate type ids by type name so that the same `T` gets the same id + // even when GetTypeId()'s function-local static is duplicated across DSOs. + ABSL_CONST_INIT static absl::Mutex mu(absl::kConstInit); + static absl::NoDestructor> + type_ids; + + absl::MutexLock lock(mu); + if (auto it = type_ids->find(type_name); it != type_ids->end()) { + return it->second; + } + size_t type_id = GetTypeInfoTableSingleton().emplace_back(type_info) + 1; DCHECK(type_id < std::numeric_limits::max()) << "Too many different AsyncValue types."; - return type_id; + type_ids->emplace(type_name, static_cast(type_id)); + return static_cast(type_id); } AsyncValue::TypeInfoTable& AsyncValue::GetTypeInfoTableSingleton() { diff --git a/third_party/xla/xla/tsl/concurrency/async_value.h b/third_party/xla/xla/tsl/concurrency/async_value.h index 1a57657bb876a6..b7e7445d0cc6f2 100644 --- a/third_party/xla/xla/tsl/concurrency/async_value.h +++ b/third_party/xla/xla/tsl/concurrency/async_value.h @@ -24,11 +24,13 @@ limitations under the License. #include #include #include +#include #include #include "absl/base/optimization.h" #include "absl/functional/any_invocable.h" #include "absl/status/status.h" +#include "absl/strings/string_view.h" #include "absl/types/span.h" #include "xla/tsl/concurrency/concurrent_vector.h" #include "xla/tsl/concurrency/executor.h" @@ -351,7 +353,13 @@ class AsyncValue { template static uint16_t CreateTypeInfoAndReturnTypeId() { return CreateTypeInfoAndReturnTypeIdImpl( - MakeTypeInfo>()); + TypeName(), MakeTypeInfo>()); + } + + // Process-stable key for `T`, used to deduplicate type ids across DSOs. + template + static absl::string_view TypeName() { + return typeid(T).name(); } std::atomic refcount_{1}; @@ -469,7 +477,8 @@ class AsyncValue { }; } - static uint16_t CreateTypeInfoAndReturnTypeIdImpl(const TypeInfo& type_info); + static uint16_t CreateTypeInfoAndReturnTypeIdImpl(absl::string_view type_name, + const TypeInfo& type_info); template T& GetConcreteValue() const; From b02620e76aebb346d8595e9a2d98da3d3b93a46f Mon Sep 17 00:00:00 2001 From: Stanislav Bardyuk Date: Wed, 30 Sep 2026 04:13:07 -0700 Subject: [PATCH 53/58] PR #47649: [XLA:GPU] Keep scatter window writes coalesced in ScatterSimplifier and the transpose folding MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Imported from GitHub PR https://github.com/openxla/xla/pull/47649 ## 📝 Summary of Changes Two changes that together restore coalesced memory writes for scatters whose update window dims come before the indexed dims (for example a vmapped `segment_sum`, where the batch dim is a window dim and the segment ids index the minor dim): - `ScatterSimplifier` gets a `reorder_operand_dims_for_coalescing` option (enabled in the GPU pipelines only). When a scatter would write strided windows under the default layouts, the operand dims are permuted so that `scatter_dims_to_operand_dims` becomes the identity mapping and the writes become contiguous, with transposes around the scatter restoring the original order. This conditionally restores the canonicalization that 087e1a5960 removed: scatters that are already coalesced keep the current no-transpose behavior, and so do variadic scatters, scatters with batching dims, and scatters whose written volume is small relative to the operand (where the transpose copies would cost more than the coalescing wins). - `TryFoldTransposeIntoScatter` in the algebraic simplifier now checks the same profitability predicate and refuses folds that would make the written windows less contiguous. Without this, the fold both undoes the ScatterSimplifier rewrite above and defeats user-side workarounds (a restoring `.transpose()` gets folded back into the scatter, recreating the strided form). Folds that improve or preserve contiguity still fire. The shared predicate `ScatterSimplifier::WriteRunLength` computes the length of the contiguous run of operand elements each scatter index writes. ## 🎯 Justification Fixes #47203 (cross-post of jax-ml/jax#39959): `segment_sum` under `vmap` regressed ~6x on the reporter's GPU between JAX 0.9.1 and 0.9.2, and its `out_axes=1` workaround regressed further in 0.10.2. Root cause chain: - Since 087e1a5960, nothing in the GPU pipeline re-orients a scatter whose window dims are major, so every scatter index writes a strided window (strided atomics). Isolated on an RTX 2070 with identical indices and updates, only the operand orientation differing: 5.11 ms coalesced vs 91.62 ms strided (18x). - Since 8a228780ee, the unconditional transpose fold recreates the strided form out of the coalesced-scatter-plus-transpose pattern, so no HLO-level workaround survives. With this PR, the issue's repro compiles back to the coalesced scatter plus one cheap restoring transpose: 85-90 ms -> ~5.5 ms on the RTX 2070 (~16x), details in the Benchmark section. ## 🚀 Kind of Contribution ⚡️ Performance Improvement / 🐛 Bug Fix ## 📊 Benchmark Issue repro as a standalone HLO (vmapped segment_sum: scatter-add of f32[1024,75960] updates into 12123 segments, hashed pseudo-random segment ids, `hlo_runner_main_gpu --num_repeats=10`, RTX 2070, sm_75): - before: 85-90 ms per execution, compiled scatter `f32[1024,12123]{1,0}` (strided window writes) - after: 5.4-5.6 ms per execution (~16x), compiled scatter `f32[12123,1024]{1,0}` (coalesced) plus a restoring transpose Isolated scatter kernel with identical indices and updates, only the operand orientation differing (via JAX on the same GPU): 5.11 ms coalesced vs 91.62 ms strided (18x). A standalone `[12123,1024] -> [1024,12123]` transpose costs 0.32 ms, so the inserted transposes are ~2% of the win. ## 🧪 Unit Tests: - `scatter_simplifier_test.cc`: `ReordersOperandDimsForCoalescing`, `ReordersOperandDimsWithInsertedWindowDims` (non-monotonic `scatter_dims_to_operand_dims`), `DoesNotReorderCoalescedScatter`, `DoesNotReorderVariadicScatter`, `DoesNotReorderWhenUpdatesAreSmall`, `DoesNotReorderOperandDimsByDefault`. - `algebraic_simplifier_test.cc`: `FoldTransposeIntoScatter` now uses a profitable example (window dims move to minor); `DoNotFoldTransposeIntoScatterWhenWritesBecomeStrided` covers the refused direction. ## 🧪 Execution Tests: Verified on a single RTX 2070 (sm_75): - `//xla/tests:scatter_test_nvgpu_any` (36 cases) and `//xla/tests:select_and_scatter_test_nvgpu_any` pass. - `//xla/service/gpu:gpu_compiler_test_nvgpu_any` and the GPU scatter emitter lit tests (`add`, `sorted_indices`, `permuted_indices`, `permuted_sorted_indices`) pass. - `run_hlo_module --platform=gpu --reference_platform=interpreter` on a small version of the repro: results match. Copybara import of the project: -- bd08ecb02fa0c0e2b705d15d56bdebfd1ac3551f by Stanislav Bardyuk : [XLA:GPU] Keep scatter window writes coalesced in ScatterSimplifier and the transpose folding Since 087e1a5960, nothing re-orients a scatter whose update window dims are major, so each scatter index writes a strided window (strided atomics, measured 18x slower than coalesced on an RTX 2070). Since 8a228780ee, the unconditional transpose-into-scatter fold recreates the strided form from coalesced-scatter-plus-transpose patterns, defeating workarounds. Add ScatterSimplifier::WriteRunLength (contiguous write-run length under default layouts), use it to gate TryFoldTransposeIntoScatter, and add a reorder_operand_dims_for_coalescing option to ScatterSimplifier (enabled on the GPU pipelines) that restores the pre-087e1a5960 operand permutation when it grows the write run and the written volume is large enough to pay for the transposes. Already-coalesced, variadic, batched, and small-update scatters keep the current no-transpose behavior. On the issue's vmapped segment_sum repro (RTX 2070), execution goes from 85-90 ms (strided f32[1024,12123] scatter) to 5.4-5.6 ms (coalesced f32[12123,1024] scatter plus a restoring transpose). Fixes #47203. Merging this change closes #47649 PiperOrigin-RevId: 990896312 --- .../simplifiers/algebraic_simplifier.cc | 8 + .../simplifiers/algebraic_simplifier_test.cc | 45 +++- third_party/xla/xla/service/BUILD | 1 + .../xla/xla/service/gpu/gpu_compiler.cc | 12 +- .../xla/xla/service/scatter_simplifier.cc | 129 +++++++++++- .../xla/xla/service/scatter_simplifier.h | 34 ++- .../xla/service/scatter_simplifier_test.cc | 199 ++++++++++++++++++ 7 files changed, 411 insertions(+), 17 deletions(-) diff --git a/third_party/xla/xla/hlo/transforms/simplifiers/algebraic_simplifier.cc b/third_party/xla/xla/hlo/transforms/simplifiers/algebraic_simplifier.cc index da7187f072fbca..7e4ed6d0be970b 100644 --- a/third_party/xla/xla/hlo/transforms/simplifiers/algebraic_simplifier.cc +++ b/third_party/xla/xla/hlo/transforms/simplifiers/algebraic_simplifier.cc @@ -10032,6 +10032,14 @@ absl::StatusOr AlgebraicSimplifierVisitor::TryFoldTransposeIntoScatter( } absl::Span permutation = transpose->dimensions(); + // Folding makes the scatter write through the transposed operand. Bail if + // that makes the written windows less contiguous than they are now: + // strided window writes do not coalesce and can cost far more than the + // transpose this rewrite saves. + if (ScatterSimplifier::WriteRunLength(scatter, permutation) < + ScatterSimplifier::WriteRunLength(scatter)) { + return false; + } std::vector inverse_permutation = InversePermutation(permutation); // Step 1 : Transpose base operand diff --git a/third_party/xla/xla/hlo/transforms/simplifiers/algebraic_simplifier_test.cc b/third_party/xla/xla/hlo/transforms/simplifiers/algebraic_simplifier_test.cc index 8af4270bb1c1ba..268e9073a943c8 100644 --- a/third_party/xla/xla/hlo/transforms/simplifiers/algebraic_simplifier_test.cc +++ b/third_party/xla/xla/hlo/transforms/simplifiers/algebraic_simplifier_test.cc @@ -15010,6 +15010,8 @@ TEST_F(AlgebraicSimplifierTest, CommuteReduceAndBroadcastUnsorted) { } TEST_F(AlgebraicSimplifierTest, FoldTransposeIntoScatter) { + // The scatter writes strided [10,1] columns; folding the transpose makes it + // write contiguous [1,10] rows, so the fold is profitable. const std::string& hlo_string = R"( HloModule m @@ -15022,12 +15024,12 @@ TEST_F(AlgebraicSimplifierTest, FoldTransposeIntoScatter) { ENTRY test { operand = f32[10, 20] parameter(0) indices = s32[5, 1] parameter(1) - updates = f32[5, 1, 20] parameter(2) + updates = f32[5, 10, 1] parameter(2) scatter = f32[10, 20] scatter(operand, indices, updates), update_window_dims={1, 2}, inserted_window_dims={}, - scatter_dims_to_operand_dims={0}, + scatter_dims_to_operand_dims={1}, index_vector_dim=1, to_apply=update_computation @@ -15042,11 +15044,11 @@ TEST_F(AlgebraicSimplifierTest, FoldTransposeIntoScatter) { constexpr absl::string_view kPattern = R"( CHECK: %[[transposed_operand:.*]] = f32[20,10]{{.*}} transpose(%[[operand:.*]]), dimensions={1,0} -CHECK: %[[transposed_updates:.*]] = f32[5,20,1]{{.*}} transpose(%[[updates:.*]]), dimensions={0,2,1} +CHECK: %[[transposed_updates:.*]] = f32[5,1,10]{{.*}} transpose(%[[updates:.*]]), dimensions={0,2,1} CHECK: ROOT %[[new_scatter:.*]] = f32[20,10]{{.*}} scatter(%[[transposed_operand]], %[[indices:.*]], %[[transposed_updates]]), CHECK-SAME: update_window_dims={1,2}, CHECK-SAME: inserted_window_dims={}, -CHECK-SAME: scatter_dims_to_operand_dims={1}, +CHECK-SAME: scatter_dims_to_operand_dims={0}, CHECK-SAME: index_vector_dim=1 )"; ASSERT_OK_AND_ASSIGN(bool matched, @@ -15054,6 +15056,41 @@ CHECK-SAME: index_vector_dim=1 EXPECT_TRUE(matched); } +TEST_F(AlgebraicSimplifierTest, + DoNotFoldTransposeIntoScatterWhenWritesBecomeStrided) { + // The scatter writes contiguous [1,20] rows; folding the transpose would + // make it write strided [20,1] columns, which does not coalesce on GPUs. + const std::string& hlo_string = R"( + HloModule m + + update_computation { + a_val = f32[] parameter(0) + b_val = f32[] parameter(1) + ROOT add = f32[] add(a_val, b_val) + } + + ENTRY test { + operand = f32[10, 20] parameter(0) + indices = s32[5, 1] parameter(1) + updates = f32[5, 1, 20] parameter(2) + + scatter = f32[10, 20] scatter(operand, indices, updates), + update_window_dims={1, 2}, + inserted_window_dims={}, + scatter_dims_to_operand_dims={0}, + index_vector_dim=1, + to_apply=update_computation + + ROOT transpose = f32[20, 10] transpose(scatter), dimensions={1, 0} + } + )"; + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_string)); + AlgebraicSimplifierOptions options = default_options_; + options.set_enable_fold_transpose_into_scatter(true); + EXPECT_THAT(AlgebraicSimplifier(options).Run(module.get()), + absl_testing::IsOkAndHolds(false)); +} + TEST_F(AlgebraicSimplifierTest, DoNotFoldTransposeIntoScatterWithInputBatchingDim) { ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(R"( diff --git a/third_party/xla/xla/service/BUILD b/third_party/xla/xla/service/BUILD index 7ed4d95dd64d5c..8f3c74e72985ea 100644 --- a/third_party/xla/xla/service/BUILD +++ b/third_party/xla/xla/service/BUILD @@ -5633,6 +5633,7 @@ cc_library( ":call_inliner", ":gather_scatter_utils", ":hlo_creation_utils", + "//xla:permutation_util", "//xla:shape_util", "//xla:util", "//xla:xla_data_proto_cc", diff --git a/third_party/xla/xla/service/gpu/gpu_compiler.cc b/third_party/xla/xla/service/gpu/gpu_compiler.cc index 012b05457db289..b256e1bb8bd482 100644 --- a/third_party/xla/xla/service/gpu/gpu_compiler.cc +++ b/third_party/xla/xla/service/gpu/gpu_compiler.cc @@ -979,7 +979,8 @@ absl::Status RunOptimizationPasses( pipeline.AddPass(); pipeline.AddPass(GatherExpander::kEliminateSimpleGathers); - pipeline.AddPass(); + pipeline.AddPass( + /*reorder_operand_dims_for_coalescing=*/true); pipeline.AddPass( ScatterExpander::kEliminateSimpleScatters); pipeline.AddPass(); @@ -1667,7 +1668,8 @@ absl::Status RunLayoutNormalizationPasses( layout_normalization_pipeline.AddPass(); // Layout normalization will create scatters that are not simplified and // also have unsorted update_window_dims. - layout_normalization_pipeline.AddPass(); + layout_normalization_pipeline.AddPass( + /*reorder_operand_dims_for_coalescing=*/true); return layout_normalization_pipeline .Run(hlo_module, {HloInstruction::kMainExecutionThread}) .status(); @@ -2181,7 +2183,8 @@ absl::Status GpuCompiler::OptimizeHloPostLayoutAssignment( gpu_version); // Layout normalization will create scatters that are not simplified and // also have unsorted update_window_dims. - pipeline.AddPass(); + pipeline.AddPass( + /*reorder_operand_dims_for_coalescing=*/true); pipeline.AddPass(); pipeline.AddPass(); pipeline.AddPass(); @@ -2269,7 +2272,8 @@ absl::Status GpuCompiler::OptimizeHloPostLayoutAssignment( // Layout normalization will create scatters that are not simplified and // also have unsorted update_window_dims. - pipeline.AddPass(); + pipeline.AddPass( + /*reorder_operand_dims_for_coalescing=*/true); // Verify the host memory space before the host offloader pass auto verifier_metadata = std::make_unique( diff --git a/third_party/xla/xla/service/scatter_simplifier.cc b/third_party/xla/xla/service/scatter_simplifier.cc index 24c03ccf42d437..bb0f7d57f627b4 100644 --- a/third_party/xla/xla/service/scatter_simplifier.cc +++ b/third_party/xla/xla/service/scatter_simplifier.cc @@ -17,6 +17,7 @@ limitations under the License. #include #include +#include #include #include "absl/algorithm/container.h" @@ -26,10 +27,12 @@ limitations under the License. #include "xla/hlo/ir/hlo_casting_utils.h" #include "xla/hlo/ir/hlo_instruction.h" #include "xla/hlo/ir/hlo_instructions.h" +#include "xla/permutation_util.h" #include "xla/service/call_inliner.h" #include "xla/service/gather_scatter_utils.h" #include "xla/service/hlo_creation_utils.h" #include "xla/shape.h" +#include "xla/shape_util.h" #include "xla/util.h" #include "xla/xla_data.pb.h" @@ -75,11 +78,25 @@ absl::StatusOr FlattenAndTransposeUpdates( return updates; } +std::vector MakeUpdatePermutation( + const std::vector& operand_permutation) { + // For the updates, we need to add the scatter dimension to the permutation. + std::vector update_permutation; + update_permutation.reserve(operand_permutation.size() + 1); + // After FlattenAndTransposeUpdates, the single scatter dimension is leading, + // keep it that way. + update_permutation.push_back(0); + for (int64_t dim : operand_permutation) { + update_permutation.push_back(dim + 1); + } + return update_permutation; +} // Transforms the scatter_updates field of scatter. scatter_indices_size is the // size of the scatter dimension in scatter_indices. absl::StatusOr> TransformScatterUpdates( HloScatterInstruction* scatter, + const std::vector& update_permutation, int64_t scatter_indices_size) { std::vector scatter_updates; const auto& attrs = scatter->scatter_dimension_numbers(); @@ -90,7 +107,7 @@ absl::StatusOr> TransformScatterUpdates( update, attrs.update_window_dims(), attrs.inserted_window_dims(), scatter_indices_size)); } - return scatter_updates; + return MaybeTranspose(scatter_updates, update_permutation); } ScatterDimensionNumbers MakeScatterDimensionNumbers( @@ -112,6 +129,36 @@ ScatterDimensionNumbers MakeScatterDimensionNumbers( return dim_numbers; } +// Returns true if permuting the operand dimensions so that +// scatter_dims_to_operand_dims maps to the leading dimensions makes the +// scatter's window writes more contiguous. +bool ShouldReorderOperandDims(const HloScatterInstruction* scatter) { + const auto& attrs = scatter->scatter_dimension_numbers(); + if (!attrs.input_batching_dims().empty() || + !attrs.scatter_indices_batching_dims().empty() || + scatter->scatter_operand_count() != 1) { + return false; + } + const Shape& operand_shape = scatter->scatter_operands().front()->shape(); + const Shape& updates_shape = scatter->scatter_updates().front()->shape(); + // The reorder wraps the scatter in two operand-sized transposes. Only + // reorder when the written volume is large enough for the coalescing win + // (strided writes measured an order of magnitude slower) to pay for the + // copies. + constexpr int64_t kMaxOperandToUpdatesRatio = 4; + if (ShapeUtil::ElementsIn(updates_shape) * kMaxOperandToUpdatesRatio < + ShapeUtil::ElementsIn(operand_shape)) { + return false; + } + const int64_t operand_rank = operand_shape.dimensions().size(); + std::vector permutation = + MakeOperandStartIndexPermutations(attrs.scatter_dims_to_operand_dims(), + operand_rank) + .first; + return ScatterSimplifier::WriteRunLength(scatter, permutation) > + ScatterSimplifier::WriteRunLength(scatter); +} + } // namespace absl::StatusOr ScatterSimplifier::ExpandInstruction( @@ -146,17 +193,45 @@ absl::StatusOr ScatterSimplifier::ExpandInstruction( return map[call_op]; } + // When requested, permute the operand dimensions so that + // scatter_dims_to_operand_dims maps to the leading dimensions, but only + // when that makes the written windows more contiguous. Strided window + // writes do not coalesce on GPUs, which costs more than the transposes + // this inserts. + std::vector operand_permutation(operand_rank); + std::vector operand_permutation_inverse(operand_rank); + absl::c_iota(operand_permutation, 0); + absl::c_iota(operand_permutation_inverse, 0); + if (reorder_operand_dims_for_coalescing_ && + ShouldReorderOperandDims(scatter)) { + auto permutations = MakeOperandStartIndexPermutations( + attrs.scatter_dims_to_operand_dims(), operand_rank); + operand_permutation = std::move(permutations.first); + operand_permutation_inverse = std::move(permutations.second); + } + auto update_permutation = MakeUpdatePermutation(operand_permutation); + ABSL_ASSIGN_OR_RETURN(auto* scatter_indices, TransformStartIndices(scatter->scatter_indices(), attrs.index_vector_dim())); ABSL_ASSIGN_OR_RETURN( auto scatter_updates, - TransformScatterUpdates(scatter, scatter_indices->shape().dimensions(0))); + TransformScatterUpdates(scatter, update_permutation, + scatter_indices->shape().dimensions(0))); + ABSL_ASSIGN_OR_RETURN( + auto scatter_operands, + MaybeTranspose(scatter->scatter_operands(), operand_permutation)); + // Map scatter_dims_to_operand_dims into the permuted operand. + std::vector scatter_dims_to_operand_dims; + scatter_dims_to_operand_dims.reserve( + attrs.scatter_dims_to_operand_dims().size()); + for (int64_t dim : attrs.scatter_dims_to_operand_dims()) { + scatter_dims_to_operand_dims.push_back(operand_permutation_inverse[dim]); + } auto dim_numbers = MakeScatterDimensionNumbers( operand_rank, attrs.scatter_dims_to_operand_dims().size(), - attrs.scatter_dims_to_operand_dims()); - const auto& scatter_operands = scatter->scatter_operands(); + scatter_dims_to_operand_dims); Shape output_shape; if (scatter_operands.size() == 1) { output_shape = scatter_operands.front()->shape(); @@ -174,7 +249,13 @@ absl::StatusOr ScatterSimplifier::ExpandInstruction( // TODO(unknown): Is this still correct? scatter->indices_are_sorted(), scatter->unique_indices())); + if (IsIdentityPermutation(operand_permutation)) { return result; + } + + // ShouldReorderOperandDims only allows non-variadic scatters, so the + // result is a single array. + return MaybeTranspose(result, operand_permutation_inverse); } bool ScatterSimplifier::IsSimplifiedScatter( @@ -201,9 +282,47 @@ bool ScatterSimplifier::IsSimplifiedScatter( dims.inserted_window_dims().empty(); } +int64_t ScatterSimplifier::WriteRunLength( + const HloScatterInstruction* scatter, + absl::Span permutation) { + const auto& dims = scatter->scatter_dimension_numbers(); + const Shape& operand_shape = scatter->scatter_operands().front()->shape(); + const Shape& updates_shape = scatter->scatter_updates().front()->shape(); + const int64_t rank = operand_shape.dimensions().size(); + + // Window slice size along each operand dimension. Inserted window dims and + // operand batching dims write a single element. + std::vector window_sizes(rank, 1); + for (int64_t dim = 0, window_dim = 0; dim < rank; ++dim) { + if (absl::c_linear_search(dims.inserted_window_dims(), dim) || + absl::c_linear_search(dims.input_batching_dims(), dim)) { + continue; + } + window_sizes[dim] = + updates_shape.dimensions(dims.update_window_dims(window_dim++)); + } + + // Walk from the minor-most dimension outwards while the writes stay + // contiguous, i.e. while the window covers each dimension fully. + int64_t run_length = 1; + for (int64_t i = rank - 1; i >= 0; --i) { + int64_t dim = permutation.empty() ? i : permutation[i]; + run_length *= window_sizes[dim]; + if (window_sizes[dim] != operand_shape.dimensions(dim)) { + break; + } + } + return run_length; +} + bool ScatterSimplifier::InstructionMatchesPattern(HloInstruction* inst) { auto* scatter = DynCast(inst); - return scatter && !IsSimplifiedScatter(scatter); + if (scatter == nullptr) { + return false; + } + return !IsSimplifiedScatter(scatter) || + (reorder_operand_dims_for_coalescing_ && + ShouldReorderOperandDims(scatter)); } } // namespace xla diff --git a/third_party/xla/xla/service/scatter_simplifier.h b/third_party/xla/xla/service/scatter_simplifier.h index cd63697df6648d..97af313372bbe0 100644 --- a/third_party/xla/xla/service/scatter_simplifier.h +++ b/third_party/xla/xla/service/scatter_simplifier.h @@ -16,6 +16,9 @@ limitations under the License. #ifndef XLA_SERVICE_SCATTER_SIMPLIFIER_H_ #define XLA_SERVICE_SCATTER_SIMPLIFIER_H_ +#include + +#include "absl/types/span.h" #include "xla/hlo/ir/hlo_instructions.h" #include "xla/hlo/transforms/expanders/op_expander_pass.h" @@ -25,10 +28,12 @@ namespace xla { // reshapes and a simpler scatter. // // It implements the first two steps of the algorithm described in -// ScatterExpander::ExpandInstruction (scatter_expander.cc). Additionally, it -// transposes updates and operands to transform scatter_dims_to_operand_dims -// into the identity mapping. This is different from the algorithm in -// ScatterExpander, which instead applies the mapping in scatter_indices. +// ScatterExpander::ExpandInstruction (scatter_expander.cc). With +// `reorder_operand_dims_for_coalescing` set, it additionally transposes +// updates and operands to transform scatter_dims_to_operand_dims into the +// identity mapping when that makes the written windows contiguous. This is +// different from the algorithm in ScatterExpander, which instead applies the +// mapping in scatter_indices. // // The semantics of the output "simple" scatter are indeed simpler than that // of a general scatter. If a scatter is simple (see IsSimplifiedScatter() for @@ -58,15 +63,36 @@ namespace xla { // Examples of simple scatter can be found in scatter_simplifier_test.cc. class ScatterSimplifier : public OpExpanderPass { public: + // If `reorder_operand_dims_for_coalescing` is set, scatters whose window + // writes would be strided under default layouts are additionally rewritten + // to permute the operand dimensions so that scatter_dims_to_operand_dims + // becomes the identity mapping and the written windows become more + // contiguous. The scatter is wrapped in transposes that restore the + // original dimension order. + explicit ScatterSimplifier(bool reorder_operand_dims_for_coalescing = false) + : reorder_operand_dims_for_coalescing_( + reorder_operand_dims_for_coalescing) {} + absl::string_view name() const override { return "scatter_simplifier"; } static bool IsSimplifiedScatter(const HloScatterInstruction* scatter); + // Returns the length of the contiguous run of operand elements that each + // scatter index writes, assuming default (descending) layouts, if the + // operand dimensions were reordered by `permutation`. An empty + // `permutation` keeps the current dimension order. Longer runs coalesce + // better on GPUs. + static int64_t WriteRunLength(const HloScatterInstruction* scatter, + absl::Span permutation = {}); + protected: bool InstructionMatchesPattern(HloInstruction* inst) override; absl::StatusOr ExpandInstruction( HloInstruction* inst) override; + + private: + const bool reorder_operand_dims_for_coalescing_; }; } // namespace xla diff --git a/third_party/xla/xla/service/scatter_simplifier_test.cc b/third_party/xla/xla/service/scatter_simplifier_test.cc index 8d6b7808c59408..f71a907ddd0d5d 100644 --- a/third_party/xla/xla/service/scatter_simplifier_test.cc +++ b/third_party/xla/xla/service/scatter_simplifier_test.cc @@ -344,6 +344,205 @@ TEST_F(ScatterSimplifierTest, VariadicScatterIntoScalar) { )"); } +TEST_F(ScatterSimplifierTest, ReordersOperandDimsForCoalescing) { + // The scatter writes strided [4,1] columns of the operand. With + // reorder_operand_dims_for_coalescing, the operand dims are permuted so + // that the writes become contiguous [1,4] rows, with transposes restoring + // the original order. + constexpr absl::string_view kModuleStr = R"( + HloModule scatter_simplifier + + scatter_computation { + p0 = f32[] parameter(0) + p1 = f32[] parameter(1) + ROOT add = f32[] add(p0, p1) + } + + ENTRY kernel_entry { + operand = f32[4,10] parameter(0) + indices = s32[5,1] parameter(1) + update = f32[5,4,1] parameter(2) + ROOT scatter = f32[4,10] scatter(operand, indices, update), + to_apply=scatter_computation, + update_window_dims={1,2}, + inserted_window_dims={}, + scatter_dims_to_operand_dims={1}, + index_vector_dim=1 + })"; + + RunAndFilecheckHloRewrite( + kModuleStr, + ScatterSimplifier(/*reorder_operand_dims_for_coalescing=*/true), R"( + CHECK: %[[OPERAND:.*]] = f32[10,4]{1,0} transpose(%operand), dimensions={1,0} + CHECK: %[[UPDATE:.*]] = f32[5,1,4]{2,1,0} transpose(%update), dimensions={0,2,1} + CHECK: %[[SCATTER:.*]] = f32[10,4]{1,0} scatter(%[[OPERAND]], %indices, %[[UPDATE]]), + CHECK-SAME: update_window_dims={1,2}, + CHECK-SAME: inserted_window_dims={}, + CHECK-SAME: scatter_dims_to_operand_dims={0}, + CHECK-SAME: index_vector_dim=1, + CHECK: ROOT %{{.*}} = f32[4,10]{1,0} transpose(%[[SCATTER]]), dimensions={1,0} + )"); +} + +TEST_F(ScatterSimplifierTest, DoesNotReorderCoalescedScatter) { + // The scatter already writes contiguous [1,4] rows; no reason to permute. + constexpr absl::string_view kModuleStr = R"( + HloModule scatter_simplifier + + scatter_computation { + p0 = f32[] parameter(0) + p1 = f32[] parameter(1) + ROOT add = f32[] add(p0, p1) + } + + ENTRY kernel_entry { + operand = f32[10,4] parameter(0) + indices = s32[5,1] parameter(1) + update = f32[5,1,4] parameter(2) + ROOT scatter = f32[10,4] scatter(operand, indices, update), + to_apply=scatter_computation, + update_window_dims={1,2}, + inserted_window_dims={}, + scatter_dims_to_operand_dims={0}, + index_vector_dim=1 + })"; + + RunAndFilecheckHloRewrite( + kModuleStr, + ScatterSimplifier(/*reorder_operand_dims_for_coalescing=*/true), + std::nullopt); +} + +TEST_F(ScatterSimplifierTest, ReordersOperandDimsWithInsertedWindowDims) { + // A non-monotonic scatter_dims_to_operand_dims with an inserted window dim: + // the reorder must remap the dimension numbers through the permutation. + constexpr absl::string_view kModuleStr = R"( + HloModule scatter_simplifier + + scatter_computation { + p0 = f32[] parameter(0) + p1 = f32[] parameter(1) + ROOT add = f32[] add(p0, p1) + } + + ENTRY kernel_entry { + operand = f32[4,6,10] parameter(0) + indices = s32[50,2] parameter(1) + update = f32[50,4,1] parameter(2) + ROOT scatter = f32[4,6,10] scatter(operand, indices, update), + to_apply=scatter_computation, + update_window_dims={1,2}, + inserted_window_dims={1}, + scatter_dims_to_operand_dims={2,1}, + index_vector_dim=1 + })"; + + RunAndFilecheckHloRewrite( + kModuleStr, + ScatterSimplifier(/*reorder_operand_dims_for_coalescing=*/true), R"( + CHECK: %[[OPERAND:.*]] = f32[10,6,4]{2,1,0} transpose(%operand), dimensions={2,1,0} + CHECK: %[[SCATTER:.*]] = f32[10,6,4]{2,1,0} scatter(%[[OPERAND]], %indices, + CHECK-SAME: update_window_dims={1,2,3}, + CHECK-SAME: inserted_window_dims={}, + CHECK-SAME: scatter_dims_to_operand_dims={0,1}, + CHECK-SAME: index_vector_dim=1, + CHECK: ROOT %{{.*}} = f32[4,6,10]{2,1,0} transpose(%[[SCATTER]]), dimensions={2,1,0} + )"); +} + +TEST_F(ScatterSimplifierTest, DoesNotReorderVariadicScatter) { + // Variadic scatters are not reordered. + constexpr absl::string_view kModuleStr = R"( + HloModule scatter_simplifier + + scatter_computation { + p0 = f32[] parameter(0) + p1 = f32[] parameter(1) + p2 = f32[] parameter(2) + p3 = f32[] parameter(3) + add0 = f32[] add(p0, p2) + add1 = f32[] add(p1, p3) + ROOT tuple = tuple(add0, add1) + } + + ENTRY kernel_entry { + operand0 = f32[4,10] parameter(0) + operand1 = f32[4,10] parameter(1) + indices = s32[5,1] parameter(2) + update0 = f32[5,4,1] parameter(3) + update1 = f32[5,4,1] parameter(4) + ROOT scatter = (f32[4,10], f32[4,10]) scatter(operand0, operand1, + indices, update0, update1), + to_apply=scatter_computation, + update_window_dims={1,2}, + inserted_window_dims={}, + scatter_dims_to_operand_dims={1}, + index_vector_dim=1 + })"; + + RunAndFilecheckHloRewrite( + kModuleStr, + ScatterSimplifier(/*reorder_operand_dims_for_coalescing=*/true), + std::nullopt); +} + +TEST_F(ScatterSimplifierTest, DoesNotReorderWhenUpdatesAreSmall) { + // The written volume is tiny compared to the operand: two operand-sized + // transposes would cost more than the coalescing wins. + constexpr absl::string_view kModuleStr = R"( + HloModule scatter_simplifier + + scatter_computation { + p0 = f32[] parameter(0) + p1 = f32[] parameter(1) + ROOT add = f32[] add(p0, p1) + } + + ENTRY kernel_entry { + operand = f32[2,262144] parameter(0) + indices = s32[3,1] parameter(1) + update = f32[3,2,1] parameter(2) + ROOT scatter = f32[2,262144] scatter(operand, indices, update), + to_apply=scatter_computation, + update_window_dims={1,2}, + inserted_window_dims={}, + scatter_dims_to_operand_dims={1}, + index_vector_dim=1 + })"; + + RunAndFilecheckHloRewrite( + kModuleStr, + ScatterSimplifier(/*reorder_operand_dims_for_coalescing=*/true), + std::nullopt); +} + +TEST_F(ScatterSimplifierTest, DoesNotReorderOperandDimsByDefault) { + // Without reorder_operand_dims_for_coalescing, a simplified scatter with + // permuted scatter_dims_to_operand_dims is left alone. + constexpr absl::string_view kModuleStr = R"( + HloModule scatter_simplifier + + scatter_computation { + p0 = f32[] parameter(0) + p1 = f32[] parameter(1) + ROOT add = f32[] add(p0, p1) + } + + ENTRY kernel_entry { + operand = f32[4,10] parameter(0) + indices = s32[5,1] parameter(1) + update = f32[5,4,1] parameter(2) + ROOT scatter = f32[4,10] scatter(operand, indices, update), + to_apply=scatter_computation, + update_window_dims={1,2}, + inserted_window_dims={}, + scatter_dims_to_operand_dims={1}, + index_vector_dim=1 + })"; + + RunAndFilecheckHloRewrite(kModuleStr, ScatterSimplifier(), std::nullopt); +} + class SimpleScatterExampleTest : public HloHardwareIndependentTestBase {}; TEST_F(SimpleScatterExampleTest, 1x1d) { From 717f95838f63fc5725767cb1742e6896bd9f0e3a Mon Sep 17 00:00:00 2001 From: Alexandros Theodoridis Date: Wed, 30 Sep 2026 04:20:58 -0700 Subject: [PATCH 54/58] PR #49767: [ROCm] Restrict usage of system env variables for rbe builds in bzl MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Imported from GitHub PR https://github.com/openxla/xla/pull/49767 📝 Summary of Changes Restrict rocm rbe builds to use only env variables passed by the config 🎯 Justification To achieve fully reproducible builds and better cache hits we would need a control over the env variables used during the build hence we restrict any externally set env variables. 🚀 Kind of Contribution Please remove what does not apply: ♻️ Cleanup, 📊 Benchmark (for Performance Improvements) Not relevant 🧪 Unit Tests: CI 🧪 Execution Tests: CI Copybara import of the project: -- 93e63d21e10002a3547a67deb4851bab0ef9d710 by Alexandros Theodoridis : Restrict usage of system env variables for rbe builds in bzl Merging this change closes #49767 PiperOrigin-RevId: 990899121 --- third_party/xla/build_tools/rocm/rocm_xla.bazelrc | 2 ++ 1 file changed, 2 insertions(+) diff --git a/third_party/xla/build_tools/rocm/rocm_xla.bazelrc b/third_party/xla/build_tools/rocm/rocm_xla.bazelrc index 546eed927c42f8..7dfbdf49341c26 100644 --- a/third_party/xla/build_tools/rocm/rocm_xla.bazelrc +++ b/third_party/xla/build_tools/rocm/rocm_xla.bazelrc @@ -16,6 +16,8 @@ build:rocm_rbe --tls_client_key="/data/ci-cert.key" build:rocm_rbe --spawn_strategy=remote,local build:rocm_rbe --grpc_keepalive_time=30s build:rocm_rbe --remote_cache_compression +build:rocm_rbe --incompatible_strict_action_env +build:rocm_rbe --experimental_strict_repo_env test:rocm_rbe --jobs=30 test:rocm_rbe --remote_executor=grpcs://wardite.cluster.engflow.com From 54a6c3a7af61b0ca0bad1724260018297c9ae6ff Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Eusebio=20Dur=C3=A1n=20Monta=C3=B1a?= Date: Wed, 30 Sep 2026 04:23:55 -0700 Subject: [PATCH 55/58] Skip NVTX-only annotation work when no profiler domain is attached. On load, `GpuExecutable` eagerly constructs `ModuleAnnotations` for both XProf and NVTX. When no NVTX profiler is attached (`DefaultProfilerDomain() == nullptr`), `RegisterString` returns a null handle and discards its input. Guard the NVTX-only work behind `DefaultProfilerDomain() != nullptr`: - Module-level stack-frame prefix extraction (`GetLongestSourceLocationPrefix`). - Per-instruction `Basic` payload formatting (`InstructionAsString`, `FormatSourceLocations`, and `CalledInstructionsAsString`). PiperOrigin-RevId: 990900227 --- .../xla/xla/backends/gpu/runtime/BUILD | 1 + .../xla/backends/gpu/runtime/annotation.cc | 27 ++++++-- .../backends/gpu/runtime/annotation_test.cc | 62 +++++++++++++++++++ 3 files changed, 84 insertions(+), 6 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/runtime/BUILD b/third_party/xla/xla/backends/gpu/runtime/BUILD index 32a2b62df23868..6584bc623be23c 100644 --- a/third_party/xla/xla/backends/gpu/runtime/BUILD +++ b/third_party/xla/xla/backends/gpu/runtime/BUILD @@ -95,6 +95,7 @@ xla_cc_test( "//xla/hlo/parser:hlo_parser", "@com_google_absl//absl/strings", "@com_google_googletest//:gtest_main", + "@tsl//tsl/profiler/lib:nvtx_utils", ], ) diff --git a/third_party/xla/xla/backends/gpu/runtime/annotation.cc b/third_party/xla/xla/backends/gpu/runtime/annotation.cc index 544bc25c02b8f7..91bc55b9dbf9e7 100644 --- a/third_party/xla/xla/backends/gpu/runtime/annotation.cc +++ b/third_party/xla/xla/backends/gpu/runtime/annotation.cc @@ -72,6 +72,14 @@ StringHandle RegisterString(const std::string& str) { return {}; } +template +StringHandle RegisterLazyString(F&& f) { + if (auto domain = tsl::profiler::DefaultProfilerDomain(); domain) { + return tsl::profiler::RegisterString(domain, std::forward(f)()); + } + return {}; +} + StringHandle RegisterOptionalString(const std::string& str) { return str.empty() ? nullptr : RegisterString(str); } @@ -426,8 +434,10 @@ ModuleAnnotation::ModuleAnnotation(const HloModule& mod) common_src_locations_(nullptr), module_id_(mod.unique_id()), common_stack_frames_(0) { - std::tie(common_src_locations_, common_stack_frames_) = - GetLongestSourceLocationPrefix(mod); + if (tsl::profiler::DefaultProfilerDomain() != nullptr) { + std::tie(common_src_locations_, common_stack_frames_) = + GetLongestSourceLocationPrefix(mod); + } } #if GOOGLE_CUDA @@ -595,11 +605,16 @@ InstructionAnnotation::InstructionAnnotation( : nvtx_name_str_(MakeInstructionTitle( module_annotation.longest_op_name_prefix(), inst)), nvtx_name_(RegisterString(nvtx_name_str_)) { + // Register these string lazily since they are expensive to produce and + // won't be used if there's no registered profiler. payload_ = Basic{ - RegisterString(InstructionAsString(inst)), - RegisterString( - FormatSourceLocations(inst, module_annotation.common_stack_frames())), - RegisterString("\n" + CalledInstructionsAsString(inst)), + RegisterLazyString([&] { return InstructionAsString(inst); }), + RegisterLazyString([&] { + return FormatSourceLocations(inst, + module_annotation.common_stack_frames()); + }), + RegisterLazyString( + [&] { return "\n" + CalledInstructionsAsString(inst); }), }; if (annotation_level < TraceAnnotationLevel::kDetailed) { xprof_name_str_ = nvtx_name_str_; diff --git a/third_party/xla/xla/backends/gpu/runtime/annotation_test.cc b/third_party/xla/xla/backends/gpu/runtime/annotation_test.cc index 728fddf0ff5529..d64593fefc6e23 100644 --- a/third_party/xla/xla/backends/gpu/runtime/annotation_test.cc +++ b/third_party/xla/xla/backends/gpu/runtime/annotation_test.cc @@ -23,6 +23,7 @@ limitations under the License. #include "xla/hlo/ir/hlo_instruction.h" #include "xla/hlo/ir/hlo_module.h" #include "xla/hlo/parser/hlo_parser.h" +#include "tsl/profiler/lib/nvtx_utils.h" namespace xla::gpu { namespace { @@ -164,5 +165,66 @@ TEST(AnnotationTest, AsyncCollectiveMetadata) { EXPECT_THAT(xprof_name, HasSubstr("is_spmd_generated=1")); } +TEST(AnnotationTest, ModuleAnnotationsWithStackFrames) { + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnUnverifiedModule(R"( + HloModule test + + FileNames + 1 "model.py" + + FunctionNames + 1 "forward" + 2 "layer" + + FileLocations + 1 {file_name_id=1 function_name_id=1 line=10 end_line=12 column=4 end_column=20} + 2 {file_name_id=1 function_name_id=2 line=42 end_line=42 column=8 end_column=30} + + StackFrames + 1 {file_location_id=1 parent_frame_id=1} + 2 {file_location_id=2 parent_frame_id=2} + + fused_computation { + p0 = f32[4] parameter(0) + p1 = f32[4] parameter(1) + ROOT add = f32[4] add(p0, p1), + metadata={op_type="add", op_name="jit(forward)/layer/add", stack_frame_id=2} + } + + ENTRY main { + a = f32[4] parameter(0) + b = f32[4] parameter(1) + ROOT fusion = f32[4] fusion(a, b), kind=kLoop, calls=fused_computation, + metadata={op_type="fusion", op_name="jit(forward)/layer/fusion", stack_frame_id=2} + } + )")); + + ModuleAnnotations basic_annotations(*module, TraceAnnotationLevel::kBasic); + EXPECT_EQ(basic_annotations.top_level.common_stack_frames(), + tsl::profiler::DefaultProfilerDomain() == nullptr ? 0 : 1); + EXPECT_EQ(basic_annotations.instructions.size(), module->instruction_count()); + const InstructionAnnotation& basic_fusion = + basic_annotations.instructions.at("fusion"); + EXPECT_FALSE(basic_fusion.has_detailed_annotations()); + EXPECT_EQ(basic_fusion.nvtx_name(), basic_fusion.xprof_name()); + EXPECT_THAT(basic_fusion.xprof_name(), HasSubstr("hlo_op=fusion")); + + ModuleAnnotations detailed_annotations(*module, + TraceAnnotationLevel::kDetailed); + EXPECT_EQ(detailed_annotations.instructions.size(), + module->instruction_count()); + const InstructionAnnotation& detailed_fusion = + detailed_annotations.instructions.at("fusion"); + EXPECT_TRUE(detailed_fusion.has_detailed_annotations()); + EXPECT_FALSE(detailed_fusion.is_collective_annotation()); + EXPECT_EQ(detailed_fusion.nvtx_name(), basic_fusion.nvtx_name()); + EXPECT_THAT(detailed_fusion.xprof_name(), HasSubstr("op_type=fusion")); + EXPECT_THAT(detailed_fusion.xprof_name(), + HasSubstr("op_name=jit(forward)/layer/fusion")); + EXPECT_THAT(detailed_fusion.xprof_name(), HasSubstr("source_file=model.py")); + EXPECT_THAT(detailed_fusion.xprof_name(), HasSubstr("source_line=42")); + EXPECT_THAT(detailed_fusion.xprof_name(), HasSubstr("shape=f32[4]")); +} + } // namespace } // namespace xla::gpu From 4706a67c5c845e96ae6e27da03bac88aeb6efa2f Mon Sep 17 00:00:00 2001 From: Mohammed Anany Date: Wed, 30 Sep 2026 05:01:24 -0700 Subject: [PATCH 56/58] Removing deprecated macros in XLA:GPU E2E tests and autotuner PiperOrigin-RevId: 990914610 --- .../xla/xla/backends/gpu/autotuner/BUILD | 10 - .../gpu/autotuner/block_level_emitter_test.cc | 60 ++- .../backends/gpu/autotuner/cublaslt_test.cc | 11 +- .../backends/gpu/autotuner/factory_test.cc | 4 +- .../gpu/autotuner/gpu_profiler_test.cc | 3 +- .../backends/gpu/autotuner/hipblaslt_test.cc | 54 +-- .../gpu/autotuner/legacy_cache_test.cc | 40 +- .../xla/backends/gpu/autotuner/miopen_test.cc | 34 +- .../autotuner/mx_scaled_dot_execution_test.cc | 9 +- .../xla/backends/gpu/host_offloading/BUILD | 1 - .../gpu_host_offloading_allocator_test.cc | 10 +- .../xla/backends/gpu/libraries/cutedsl/BUILD | 1 - .../gpu/libraries/cutedsl/ffi_test.cc | 4 +- .../xla/xla/backends/gpu/profiler/BUILD | 1 - .../gpu/profiler/kernel_name_tracer_test.cc | 41 +- .../xla/xla/backends/gpu/target_config/BUILD | 4 - .../target_config/cudnn_device_props_test.cc | 12 +- .../embedded_target_config_test.cc | 13 +- .../gpu/target_config/target_config_test.cc | 3 +- third_party/xla/xla/backends/gpu/tests/BUILD | 30 +- .../backends/gpu/tests/all_reduce_e2e_test.cc | 67 ++- .../gpu/tests/async_command_buffer_test.cc | 7 +- .../gpu/tests/async_kernel_launch_test.cc | 7 +- .../gpu/tests/collective_ops_e2e_test.cc | 432 +++++++++--------- .../gpu/tests/collective_ops_ffi_test.cc | 93 ++-- ...llective_ops_sharded_unsharded_e2e_test.cc | 12 +- .../collective_pipeline_parallelism_test.cc | 105 ++--- .../xla/backends/gpu/tests/gpu_atomic_test.cc | 12 +- .../gpu/tests/gpu_spmd_e2e_compile_test.cc | 3 +- .../gpu/tests/gpu_triton_custom_call_test.cc | 3 +- .../gpu/tests/multioutput_fusion_test.cc | 22 +- .../gpu/tests/nccl_group_execution_test.cc | 14 +- .../gpu/tests/nop_custom_call_test.cc | 4 +- .../backends/gpu/tests/p2p_ops_e2e_test.cc | 10 +- .../xla/backends/gpu/tests/ptx_kernel_test.cc | 22 +- .../xla/backends/gpu/tests/ragged_dot_test.cc | 26 +- .../gpu/tests/replicated_io_feed_test.cc | 15 +- .../gpu/tests/simple_optimization_test.cc | 13 +- .../xla/backends/gpu/tests/sorting_test.cc | 17 +- 39 files changed, 573 insertions(+), 656 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/autotuner/BUILD b/third_party/xla/xla/backends/gpu/autotuner/BUILD index 9309c02d17b1b3..9ba454873956d6 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/BUILD +++ b/third_party/xla/xla/backends/gpu/autotuner/BUILD @@ -137,7 +137,6 @@ xla_test( "//xla/stream_executor:platform", "//xla/stream_executor:stream_executor_h", "//xla/tsl/platform:env", - "//xla/tsl/platform:statusor", "//xla/tsl/util/proto:proto_matchers", "@com_google_absl//absl/log:check", "@com_google_absl//absl/status:status_matchers", @@ -217,7 +216,6 @@ xla_test( "//xla/stream_executor:blas", "//xla/stream_executor:device_description_proto_cc", "//xla/stream_executor:stream_executor_h", - "//xla/tsl/lib/core:status_test_util", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/status:statusor", @@ -627,7 +625,6 @@ xla_test( "//xla/backends/gpu/tests:hlo_pjrt_gpu_test_base", "//xla/hlo/parser:hlo_parser", "//xla/hlo/testlib:filecheck", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/strings", "@com_google_googletest//:gtest", @@ -661,8 +658,6 @@ xla_test( "//xla/stream_executor:blas", "//xla/stream_executor:device_description_proto_cc", "//xla/stream_executor:stream_executor_h", - "//xla/tsl/lib/core:status_test_util", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/status:statusor", @@ -750,7 +745,6 @@ xla_test( "//xla/stream_executor:platform", "//xla/stream_executor:stream_executor_address_allocator", "//xla/stream_executor:stream_executor_h", - "//xla/tsl/lib/core:status_test_util", "@com_google_absl//absl/log", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_macros", @@ -935,9 +929,7 @@ xla_cc_test( "//xla/hlo/parser:hlo_parser", "//xla/stream_executor:device_description", "//xla/stream_executor/cuda:cuda_compute_capability", - "//xla/tsl/lib/core:status_test_util", "//xla/tsl/platform:env", - "//xla/tsl/platform:statusor", "//xla/tsl/protobuf:dnn_proto_cc", "@com_google_absl//absl/status", "@com_google_googletest//:gtest_main", @@ -1108,8 +1100,6 @@ xla_test( "//xla/stream_executor:stream_executor_h", "//xla/stream_executor:stream_executor_memory_allocator", "//xla/stream_executor/rocm:rocm_platform_id", - "//xla/tsl/lib/core:status_test_util", - "//xla/tsl/platform:statusor", "//xla/tsl/protobuf:dnn_proto_cc", "//xla/tsl/util/proto:proto_matchers", "@com_google_absl//absl/status", diff --git a/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter_test.cc b/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter_test.cc index e2140b84ef0129..47d5bf1f7a39cf 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/block_level_emitter_test.cc @@ -40,7 +40,6 @@ limitations under the License. #include "xla/stream_executor/platform.h" #include "xla/stream_executor/stream_executor.h" #include "xla/tsl/platform/env.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/threadpool.h" #include "xla/tsl/util/proto/proto_matchers.h" #include "xla/xla.pb.h" @@ -99,8 +98,8 @@ class TritonBlockLevelFusionEmitterBackendTest TEST_F(TritonBlockLevelFusionEmitterBackendTest, GetDefaultConfig_FromHlo) { // Parse an HLO module containing a kCustom Triton fusion with a backend // config that includes block-level tiling parameters. - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(R"( HloModule m %wrapped_transpose_computation { %param_0 = f32[16,64]{1,0} parameter(0) @@ -128,10 +127,9 @@ ENTRY %main { )")); // Call GetDefaultConfig on the root instruction (the fusion op). - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr config, - backend_.GetDefaultConfig( - *(module->entry_computation()->root_instruction()))); + ASSERT_OK_AND_ASSIGN(std::unique_ptr config, + backend_.GetDefaultConfig( + *(module->entry_computation()->root_instruction()))); // Verify that the returned config is indeed a BlockLevelFusionConfig. ASSERT_TRUE(config->has_block_level()); BlockLevelFusionConfig block_level_fusion_config = config->block_level(); @@ -152,8 +150,8 @@ ENTRY %main { // cost model, which has its own tests. TEST_F(TritonBlockLevelFusionEmitterBackendTest, GetDefaultConfig_Fallback) { // Parse an HLO module with a fusion instruction lacking any backend config. - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(R"( HloModule m %wrapped_transpose_computation { %param_0 = f32[16,1,64]{2,1,0} parameter(0) @@ -168,10 +166,9 @@ ENTRY %main { )")); // Call GetDefaultConfig on the root instruction (the fusion op). - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr config, - backend_.GetDefaultConfig( - *(module->entry_computation()->root_instruction()))); + ASSERT_OK_AND_ASSIGN(std::unique_ptr config, + backend_.GetDefaultConfig( + *(module->entry_computation()->root_instruction()))); // Verify that the returned config is indeed a BlockLevelFusionConfig. ASSERT_TRUE(config->has_block_level()); BlockLevelFusionConfig block_level_fusion_config = config->block_level(); @@ -186,8 +183,8 @@ ENTRY %main { // BlockLevelFusionConfig to a fusion instruction. TEST_F(TritonBlockLevelFusionEmitterBackendTest, ApplyConfig) { // Build and verify a simple HLO module containing a 2D transpose fusion. - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(R"( HloModule m %wrapped_transpose_computation { %param_0 = f32[16,64]{1,0} parameter(0) @@ -205,18 +202,17 @@ ENTRY %main { // Ask the backend to generate a default block-level fusion configuration // for this fusion operation. // Call GetDefaultConfig on the root instruction (the fusion op). - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr config, - backend_.GetDefaultConfig( - *(module->entry_computation()->root_instruction()))); + ASSERT_OK_AND_ASSIGN(std::unique_ptr config, + backend_.GetDefaultConfig( + *(module->entry_computation()->root_instruction()))); // Verify that the returned config is indeed a BlockLevelFusionConfig. ASSERT_TRUE(config->has_block_level()); BlockLevelFusionConfig block_level_fusion_config = config->block_level(); // Apply the generated config to the fusion instruction. EXPECT_THAT(backend_.ApplyConfig(*instr, *config), absl_testing::IsOk()); - TF_ASSERT_OK_AND_ASSIGN(GpuBackendConfig gpu_backend_config, - instr->backend_config()); + ASSERT_OK_AND_ASSIGN(GpuBackendConfig gpu_backend_config, + instr->backend_config()); // Ensure that the backend config on the instruction matches what was applied. EXPECT_THAT( gpu_backend_config.fusion_backend_config().block_level_fusion_config(), @@ -229,8 +225,8 @@ ENTRY %main { TEST_F(TritonBlockLevelFusionEmitterBackendTest, Compile) { // Parse an HLO module containing a kCustom Triton fusion with a backend // config that includes block-level tiling parameters. - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(R"( HloModule m %wrapped_transpose_computation { %param_0 = f32[16,64]{1,0} parameter(0) @@ -257,10 +253,9 @@ ENTRY %main { } )")); // Call GetDefaultConfig on the root instruction (the fusion op). - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr config, - backend_.GetDefaultConfig( - *(module->entry_computation()->root_instruction()))); + ASSERT_OK_AND_ASSIGN(std::unique_ptr config, + backend_.GetDefaultConfig( + *(module->entry_computation()->root_instruction()))); // Attempt to compile the root instruction using the retrieved backend config. absl::StatusOr> executable = backend_.Compile( *(module->entry_computation()->root_instruction()), *config); @@ -271,8 +266,8 @@ ENTRY %main { TEST_F(TritonBlockLevelFusionEmitterBackendTest, CompileThroughCostModelConfig) { // Parse an HLO module without any assigned backend config. - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(R"( HloModule m %wrapped_transpose_computation { %param_0 = f32[16,64]{1,0} parameter(0) @@ -286,10 +281,9 @@ ENTRY %main { } )")); // Call GetDefaultConfig on the root instruction (the fusion op). - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr config, - backend_.GetDefaultConfig( - *(module->entry_computation()->root_instruction()))); + ASSERT_OK_AND_ASSIGN(std::unique_ptr config, + backend_.GetDefaultConfig( + *(module->entry_computation()->root_instruction()))); // Attempt to compile the root instruction using the retrieved backend config. absl::StatusOr> executable = backend_.Compile( *(module->entry_computation()->root_instruction()), *config); diff --git a/third_party/xla/xla/backends/gpu/autotuner/cublaslt_test.cc b/third_party/xla/xla/backends/gpu/autotuner/cublaslt_test.cc index fbb13cd07a8343..afb797c263a3ad 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/cublaslt_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/cublaslt_test.cc @@ -39,7 +39,6 @@ limitations under the License. #include "xla/stream_executor/blas.h" #include "xla/stream_executor/device_description.pb.h" #include "xla/stream_executor/stream_executor.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/xla.pb.h" namespace xla { @@ -213,11 +212,11 @@ TEST_F(CublasLtBackendTest, ApplyConfig) { config.set_autotune_workspace_size(42); BackendConfig backend_config; *backend_config.mutable_gemm() = config; - TF_EXPECT_OK(backend_.ApplyConfig(*hlo_module->entry_computation() - ->root_instruction() - ->mutable_operands() - .at(0), - backend_config)); + EXPECT_OK(backend_.ApplyConfig(*hlo_module->entry_computation() + ->root_instruction() + ->mutable_operands() + .at(0), + backend_config)); EXPECT_THAT(RunFileCheck(hlo_module->ToString(), R"(CHECK: (f32[100,100]{1,0}, s8[42]{0}) custom-call CHECK: "selected_algorithm":"2")"), diff --git a/third_party/xla/xla/backends/gpu/autotuner/factory_test.cc b/third_party/xla/xla/backends/gpu/autotuner/factory_test.cc index b2de8f1f15c566..9e777cda88dca1 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/factory_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/factory_test.cc @@ -20,6 +20,7 @@ limitations under the License. #include #include +#include #include #include "absl/log/check.h" #include "absl/status/statusor.h" @@ -38,7 +39,6 @@ limitations under the License. #include "xla/stream_executor/platform_manager.h" #include "xla/stream_executor/stream_executor.h" #include "xla/stream_executor/stream_executor_address_allocator.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla.pb.h" namespace xla { @@ -106,7 +106,7 @@ TEST_P(FactoryTest, GetCodegenBackends) { (GetParam().run_on_rocm && is_rocm)) { auto& registry = stream_executor::PlatformObjectRegistry::GetGlobalRegistry(); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( const GetCodegenBackends::Type& get_codegen_backends, registry.FindObject(platform_->id())); mlir::MLIRContext mlir_context; diff --git a/third_party/xla/xla/backends/gpu/autotuner/gpu_profiler_test.cc b/third_party/xla/xla/backends/gpu/autotuner/gpu_profiler_test.cc index 3e7ddcdb7103b6..ea166fa8954143 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/gpu_profiler_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/gpu_profiler_test.cc @@ -51,7 +51,6 @@ limitations under the License. #include "xla/stream_executor/platform.h" #include "xla/stream_executor/stream_executor.h" #include "xla/stream_executor/stream_executor_address_allocator.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/xla_data.pb.h" namespace xla { @@ -279,7 +278,7 @@ TEST_P(GpuProfilerTestWithRedzonePadding, CheckInputBuffers) { auto profiler = GpuProfiler::Create(stream_exec_, options, allocator_.get()); ASSERT_OK_AND_ASSIGN(std::unique_ptr buffers, profiler->CreateInputBuffers(&mock_executable)); - TF_EXPECT_OK(profiler->CheckInputBuffers(*buffers)); + EXPECT_OK(profiler->CheckInputBuffers(*buffers)); } INSTANTIATE_TEST_SUITE_P(GpuProfilerTestWithRedzonePadding, diff --git a/third_party/xla/xla/backends/gpu/autotuner/hipblaslt_test.cc b/third_party/xla/xla/backends/gpu/autotuner/hipblaslt_test.cc index ccad178efb4d90..d0abbd8ffb05a3 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/hipblaslt_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/hipblaslt_test.cc @@ -36,8 +36,6 @@ limitations under the License. #include "xla/stream_executor/blas.h" #include "xla/stream_executor/device_description.pb.h" #include "xla/stream_executor/stream_executor.h" -#include "xla/tsl/lib/core/status_test_util.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla.pb.h" namespace xla { @@ -129,9 +127,8 @@ TEST_F(HipblasLtBackendTest, CanCreateHipblasLtBackend) { } TEST_F(HipblasLtBackendTest, GetSupportedConfigs) { - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr hlo_module, - ParseAndReturnVerifiedModule(kHipblasLtCustomCallHlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, + ParseAndReturnVerifiedModule(kHipblasLtCustomCallHlo)); absl::StatusOr>> configs = backend_.GetSupportedConfigs( @@ -155,8 +152,8 @@ TEST_F(HipblasLtBackendTest, GetSupportedConfigsReturnsErrorForDeviceless) { TEST_F(HipblasLtBackendTest, GetSupportedConfigsReturnsEmptyVectorForNonHipblasLtCustomCall) { - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, - ParseAndReturnVerifiedModule(kUnsupportedHlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, + ParseAndReturnVerifiedModule(kUnsupportedHlo)); absl::StatusOr>> configs = backend_.GetSupportedConfigs( @@ -165,9 +162,8 @@ TEST_F(HipblasLtBackendTest, } TEST_F(HipblasLtBackendTest, GetDefaultConfig) { - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr module, - ParseAndReturnVerifiedModule(kHipblasLtCustomCallHlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(kHipblasLtCustomCallHlo)); absl::StatusOr> config = backend_.GetDefaultConfig( @@ -189,8 +185,8 @@ TEST_F(HipblasLtBackendTest, GetDefaultConfigFailsWithoutAHipblasLtCustomCall) { lhs_contracting_dims={1}, rhs_contracting_dims={0} })"; - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(hlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(hlo)); absl::StatusOr> config = backend_.GetDefaultConfig( (*module->entry_computation()->root_instruction())); @@ -199,19 +195,18 @@ TEST_F(HipblasLtBackendTest, GetDefaultConfigFailsWithoutAHipblasLtCustomCall) { } TEST_F(HipblasLtBackendTest, ApplyConfig) { - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr hlo_module, - ParseAndReturnVerifiedModule(kHipblasLtCustomCallHlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, + ParseAndReturnVerifiedModule(kHipblasLtCustomCallHlo)); HipblasLtBackendConfig config; config.set_algorithm(2); config.set_autotune_workspace_size(42); BackendConfig backend_config; *backend_config.mutable_gemm() = config; - TF_EXPECT_OK(backend_.ApplyConfig(*hlo_module->entry_computation() - ->root_instruction() - ->mutable_operands() - .at(0), - backend_config)); + EXPECT_OK(backend_.ApplyConfig(*hlo_module->entry_computation() + ->root_instruction() + ->mutable_operands() + .at(0), + backend_config)); EXPECT_THAT(RunFileCheck(hlo_module->ToString(), R"(CHECK: (f32[100,100]{1,0}, s8[42]{0}) custom-call CHECK: "selected_algorithm":"2")"), @@ -219,10 +214,9 @@ TEST_F(HipblasLtBackendTest, ApplyConfig) { } TEST_F(HipblasLtBackendTest, Compile) { - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr module, - ParseAndReturnVerifiedModule(kHipblasLtCustomCallHlo)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(kHipblasLtCustomCallHlo)); + ASSERT_OK_AND_ASSIGN( std::unique_ptr config, backend_.GetDefaultConfig( *(module->entry_computation()->root_instruction()->operand(0)))); @@ -286,7 +280,7 @@ class HipblasLtScaledDotTest : public HipblasLtBackendTest { TEST_F(HipblasLtScaledDotTest, GetSupportedConfigs) { for (const char* hlo : kScaledDotHlos) { - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); HloInstruction* fusion = module->entry_computation()->root_instruction(); auto configs = backend_.GetSupportedConfigs(*fusion); EXPECT_THAT(configs, @@ -296,11 +290,11 @@ TEST_F(HipblasLtScaledDotTest, GetSupportedConfigs) { TEST_F(HipblasLtScaledDotTest, ApplyConfig) { for (const char* hlo : kScaledDotHlos) { - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); HloInstruction* fusion = module->entry_computation()->root_instruction(); - TF_ASSERT_OK_AND_ASSIGN(auto config, backend_.GetDefaultConfig(*fusion)); + ASSERT_OK_AND_ASSIGN(auto config, backend_.GetDefaultConfig(*fusion)); - TF_EXPECT_OK(backend_.ApplyConfig(*fusion, *config)); + EXPECT_OK(backend_.ApplyConfig(*fusion, *config)); EXPECT_THAT(RunFileCheck(module->ToString(), R"(CHECK: custom-call @@ -312,9 +306,9 @@ TEST_F(HipblasLtScaledDotTest, ApplyConfig) { TEST_F(HipblasLtScaledDotTest, Compile) { for (const char* hlo : kScaledDotHlos) { - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); HloInstruction* fusion = module->entry_computation()->root_instruction(); - TF_ASSERT_OK_AND_ASSIGN(auto config, backend_.GetDefaultConfig(*fusion)); + ASSERT_OK_AND_ASSIGN(auto config, backend_.GetDefaultConfig(*fusion)); auto executable = backend_.Compile(*fusion, *config); EXPECT_THAT(executable, absl_testing::IsOk()); diff --git a/third_party/xla/xla/backends/gpu/autotuner/legacy_cache_test.cc b/third_party/xla/xla/backends/gpu/autotuner/legacy_cache_test.cc index 4fe23e6fe79ac9..8f6ef7a718e218 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/legacy_cache_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/legacy_cache_test.cc @@ -33,9 +33,7 @@ limitations under the License. #include "xla/literal_util.h" #include "xla/stream_executor/cuda/cuda_compute_capability.h" #include "xla/stream_executor/device_description.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/tsl/platform/env.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/protobuf/dnn.pb.h" #include "xla/xla.pb.h" #include "tsl/platform/protobuf.h" @@ -68,7 +66,7 @@ class LegacyCacheTest : public ::testing::Test { protected: void SetUp() override { ASSERT_TRUE(tsl::Env::Default()->LocalTempFilename(&test_dir_)); - TF_ASSERT_OK(tsl::Env::Default()->CreateDir(test_dir_)); + ASSERT_OK(tsl::Env::Default()->CreateDir(test_dir_)); } void TearDown() override { @@ -154,7 +152,7 @@ TEST_F(LegacyCacheTest, InsertAndLookupTriton) { auto instr = CreateDummyInstr("hlo1"); Config config = CreateDummyTritonConfig(); - TF_ASSERT_OK(cache.Insert(instr.get(), config)); + ASSERT_OK(cache.Insert(instr.get(), config)); EXPECT_THAT(cache.Lookup(instr.get()), Optional(ConfigEq(config))); } @@ -163,7 +161,7 @@ TEST_F(LegacyCacheTest, InsertAndLookupCublas) { auto instr = CreateDummyInstr("hlo2"); Config config = CreateDummyCublasLtConfig(); - TF_ASSERT_OK(cache.Insert(instr.get(), config)); + ASSERT_OK(cache.Insert(instr.get(), config)); EXPECT_THAT(cache.Lookup(instr.get()), Optional(ConfigEq(config))); } @@ -184,11 +182,11 @@ ENTRY main { ROOT fusion.0 = f32[] fusion(p0, p1), kind=kLoop, calls=fused_computation } )"; - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnUnverifiedModule(kHLO)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnUnverifiedModule(kHLO)); auto instr = module->entry_computation()->root_instruction(); Config config = CreateDummyCublasLtFissionConfig(); - TF_ASSERT_OK(cache.Insert(instr, config)); + ASSERT_OK(cache.Insert(instr, config)); EXPECT_THAT(cache.Lookup(instr), Optional(ConfigEq(config))); } @@ -197,7 +195,7 @@ TEST_F(LegacyCacheTest, InsertAndLookupCudnn) { auto instr = CreateDummyInstr("hlo3"); Config config = CreateDummyCudnnConfig(); - TF_ASSERT_OK(cache.Insert(instr.get(), config)); + ASSERT_OK(cache.Insert(instr.get(), config)); EXPECT_THAT(cache.Lookup(instr.get()), Optional(ConfigEq(config))); } @@ -206,7 +204,7 @@ TEST_F(LegacyCacheTest, InsertAndLookupOther) { auto instr = CreateDummyInstr("hlo5"); Config config = CreateDummyBackendConfig(); - TF_ASSERT_OK(cache.Insert(instr.get(), config)); + ASSERT_OK(cache.Insert(instr.get(), config)); std::optional actual_config = cache.Lookup(instr.get()); EXPECT_THAT(actual_config, Optional(ConfigEq(config))); } @@ -218,7 +216,7 @@ TEST_F(LegacyCacheTest, PersistAcrossInstances) { // Create cache, insert, and let it save. { auto cache = LegacyCache(test_dir_, mode_, device_desc_); - TF_ASSERT_OK(cache.Insert(instr.get(), config)); + ASSERT_OK(cache.Insert(instr.get(), config)); } // Create a new cache, which should load from disk. @@ -236,7 +234,7 @@ TEST_F(LegacyCacheTest, LoadWithDifferentDevice) { { auto cache = LegacyCache(test_dir_, mode_, CreateDummyDeviceDescription("test_device")); - TF_ASSERT_OK(cache.Insert(instr.get(), config)); + ASSERT_OK(cache.Insert(instr.get(), config)); } // Create a new cache with different device, should not load the entry. @@ -252,11 +250,11 @@ TEST_F(LegacyCacheTest, OnlyInsertOncePerHlo) { auto instr = CreateDummyInstr("hlo8"); Config config = CreateDummyTritonConfig(); - TF_ASSERT_OK(cache.Insert(instr.get(), config)); + ASSERT_OK(cache.Insert(instr.get(), config)); EXPECT_THAT(cache.Lookup(instr.get()), Optional(ConfigEq(config))); Config another_config = CreateDummyCublasLtConfig(); - TF_ASSERT_OK(cache.Insert(instr.get(), another_config)); + ASSERT_OK(cache.Insert(instr.get(), another_config)); EXPECT_THAT(cache.Lookup(instr.get()), Optional(ConfigEq(config))); } @@ -265,23 +263,23 @@ TEST_F(LegacyCacheTest, SerializeAndDeserialize) { std::unique_ptr instr_1 = CreateDummyInstr("hlo9"); std::unique_ptr instr_2 = CreateDummyInstr("hlo10"); Config orig_config = CreateDummyTritonConfig(); - TF_ASSERT_OK(cache.Insert(instr_1.get(), orig_config)); - TF_ASSERT_OK(cache.Insert(instr_2.get(), orig_config)); + ASSERT_OK(cache.Insert(instr_1.get(), orig_config)); + ASSERT_OK(cache.Insert(instr_2.get(), orig_config)); // Serialize instr_1 to a string. std::vector instructions_to_serialize = { instr_1.get()}; - TF_ASSERT_OK_AND_ASSIGN(std::string serialized_cache, - cache.Serialize(instructions_to_serialize)); + ASSERT_OK_AND_ASSIGN(std::string serialized_cache, + cache.Serialize(instructions_to_serialize)); // Overwrite config for both instructions. cache.ClearCache(); Config another_config = CreateDummyCublasLtConfig(); - TF_ASSERT_OK(cache.Insert(instr_1.get(), another_config)); - TF_ASSERT_OK(cache.Insert(instr_2.get(), another_config)); + ASSERT_OK(cache.Insert(instr_1.get(), another_config)); + ASSERT_OK(cache.Insert(instr_2.get(), another_config)); // Deserialize the cache, only instr_1 should be overwritten. - TF_ASSERT_OK(cache.Deserialize(serialized_cache)); + ASSERT_OK(cache.Deserialize(serialized_cache)); EXPECT_THAT(cache.Lookup(instr_1.get()), Optional(ConfigEq(orig_config))); EXPECT_THAT(cache.Lookup(instr_2.get()), Optional(ConfigEq(another_config))); } @@ -302,7 +300,7 @@ TEST_F(LegacyCacheTest, CacheStats) { EXPECT_EQ(cache.GetCacheStats().misses, 1); // Insert and lookup hit. - TF_ASSERT_OK(cache.Insert(instr1.get(), config)); + ASSERT_OK(cache.Insert(instr1.get(), config)); EXPECT_THAT(cache.Lookup(instr1.get()), Optional(ConfigEq(config))); EXPECT_EQ(cache.GetCacheStats().hits, 1); EXPECT_EQ(cache.GetCacheStats().misses, 1); diff --git a/third_party/xla/xla/backends/gpu/autotuner/miopen_test.cc b/third_party/xla/xla/backends/gpu/autotuner/miopen_test.cc index a7e1bc0ef63232..b675f14abcc719 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/miopen_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/miopen_test.cc @@ -37,8 +37,6 @@ limitations under the License. #include "xla/stream_executor/rocm/rocm_platform_id.h" #include "xla/stream_executor/stream_executor.h" #include "xla/stream_executor/stream_executor_memory_allocator.h" -#include "xla/tsl/lib/core/status_test_util.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/protobuf/dnn.pb.h" #include "xla/tsl/util/proto/proto_matchers.h" #include "xla/xla.pb.h" @@ -110,8 +108,8 @@ TEST_F(MIOpenBackendTest, GetSupportedConfigsFromMIOpenCustomCall) { if (!IsRocm()) { GTEST_SKIP() << "Skipping test on non-ROCm platform"; } - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, - ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, + ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); absl::StatusOr>> configs = backend_.GetSupportedConfigs( (*hlo_module->entry_computation()->root_instruction()->operand(0))); @@ -124,8 +122,8 @@ TEST_F(MIOpenBackendTest, GetSupportedConfigsFromMIOpenCustomCall) { TEST_F(MIOpenBackendTest, GetSupportedConfigsReturnsErrorForDeviceless) { MIOpenBackend backend_without_stream_executor( nullptr, &debug_options_, &compiler_, &target_config_, &allocator_); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, - ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, + ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); absl::StatusOr>> configs = backend_without_stream_executor.GetSupportedConfigs( (*hlo_module->entry_computation()->root_instruction()->operand(0))); @@ -137,8 +135,8 @@ TEST_F(MIOpenBackendTest, GetDefaultConfigFromMIOpenCustomCall) { if (!IsRocm()) { GTEST_SKIP() << "Skipping test on non-ROCm platform"; } - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, - ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, + ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); absl::StatusOr> config = backend_.GetDefaultConfig( (*hlo_module->entry_computation()->root_instruction()->operand(0))); @@ -149,17 +147,17 @@ TEST_F(MIOpenBackendTest, ApplyConfigToMIOpenCustomCall) { if (!IsRocm()) { GTEST_SKIP() << "Skipping test on non-ROCm platform"; } - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, - ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, + ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); MIOpenBackendConfig config; config.set_algo_id(1); HloInstruction* instr = hlo_module->entry_computation()->root_instruction()->mutable_operand(0); BackendConfig backend_config; *backend_config.mutable_algorithm() = config; - TF_ASSERT_OK(backend_.ApplyConfig(*instr, backend_config)); - TF_ASSERT_OK_AND_ASSIGN(GpuBackendConfig gpu_config, - instr->backend_config()); + ASSERT_OK(backend_.ApplyConfig(*instr, backend_config)); + ASSERT_OK_AND_ASSIGN(GpuBackendConfig gpu_config, + instr->backend_config()); EXPECT_THAT(gpu_config.cudnn_conv_backend_config().algorithm(), EqualsProto(config)); } @@ -168,8 +166,8 @@ TEST_F(MIOpenBackendTest, ApplyConfigToMIOpenCustomCallWithWorkspace) { if (!IsRocm()) { GTEST_SKIP() << "Skipping test on non-ROCm platform"; } - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, - ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr hlo_module, + ParseAndReturnVerifiedModule(kMIOpenCustomCallHlo)); MIOpenBackendConfig config; config.set_algo_id(1); config.mutable_workspace_size()->set_value(1024); @@ -177,13 +175,13 @@ TEST_F(MIOpenBackendTest, ApplyConfigToMIOpenCustomCallWithWorkspace) { hlo_module->entry_computation()->root_instruction()->mutable_operand(0); BackendConfig backend_config; *backend_config.mutable_algorithm() = config; - TF_ASSERT_OK(backend_.ApplyConfig(*instr, backend_config)); + ASSERT_OK(backend_.ApplyConfig(*instr, backend_config)); auto* replaced_instr = hlo_module->entry_computation()->GetInstructionWithName("cudnn-conv"); - TF_ASSERT_OK_AND_ASSIGN(GpuBackendConfig gpu_config, - replaced_instr->backend_config()); + ASSERT_OK_AND_ASSIGN(GpuBackendConfig gpu_config, + replaced_instr->backend_config()); EXPECT_THAT(gpu_config.cudnn_conv_backend_config().algorithm(), EqualsProto(config)); EXPECT_EQ(replaced_instr->shape().tuple_shapes(1).dimensions(0), 1024); diff --git a/third_party/xla/xla/backends/gpu/autotuner/mx_scaled_dot_execution_test.cc b/third_party/xla/xla/backends/gpu/autotuner/mx_scaled_dot_execution_test.cc index 77c3e3e27f77be..965fbb36fe8e8b 100644 --- a/third_party/xla/xla/backends/gpu/autotuner/mx_scaled_dot_execution_test.cc +++ b/third_party/xla/xla/backends/gpu/autotuner/mx_scaled_dot_execution_test.cc @@ -24,7 +24,6 @@ limitations under the License. #include "xla/error_spec.h" #include "xla/hlo/parser/hlo_parser.h" #include "xla/hlo/testlib/filecheck.h" -#include "xla/tsl/platform/statusor.h" namespace xla::gpu { namespace { @@ -47,8 +46,8 @@ class MxScaledDotExecutionTest : public HloPjRtGpuTestBase { ref_config.mutable_debug_options() .set_xla_gpu_experimental_scaled_dot_with_triton(false); ref_config.mutable_debug_options().set_xla_gpu_enable_triton_gemm(false); - TF_ASSERT_OK_AND_ASSIGN(auto ref_optimized, - GetOptimizedModule(hlo_string, ref_config)); + ASSERT_OK_AND_ASSIGN(auto ref_optimized, + GetOptimizedModule(hlo_string, ref_config)); EXPECT_THAT( RunFileCheck(ref_optimized->ToString(), R"(CHECK: {{__cublas\$lt\$matmul|__cublas\$gemm}})"), @@ -58,8 +57,8 @@ class MxScaledDotExecutionTest : public HloPjRtGpuTestBase { test_config.mutable_debug_options() .set_xla_gpu_experimental_scaled_dot_with_triton(true); test_config.mutable_debug_options().set_xla_gpu_enable_triton_gemm(true); - TF_ASSERT_OK_AND_ASSIGN(auto test_optimized, - GetOptimizedModule(hlo_string, test_config)); + ASSERT_OK_AND_ASSIGN(auto test_optimized, + GetOptimizedModule(hlo_string, test_config)); // The autotuner may pick any of the MX-aware ROCm backends: // __triton_nested_gemm_fusion -> Triton // __cublas$lt$matmul$mx -> hipBLASLt diff --git a/third_party/xla/xla/backends/gpu/host_offloading/BUILD b/third_party/xla/xla/backends/gpu/host_offloading/BUILD index 6fba46f1f03f73..e392ddec9652f0 100644 --- a/third_party/xla/xla/backends/gpu/host_offloading/BUILD +++ b/third_party/xla/xla/backends/gpu/host_offloading/BUILD @@ -66,7 +66,6 @@ xla_test( "//xla/stream_executor:platform_manager", "//xla/stream_executor:stream", "//xla/stream_executor:stream_executor_h", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", diff --git a/third_party/xla/xla/backends/gpu/host_offloading/gpu_host_offloading_allocator_test.cc b/third_party/xla/xla/backends/gpu/host_offloading/gpu_host_offloading_allocator_test.cc index 7adac426d53893..40813912819ab2 100644 --- a/third_party/xla/xla/backends/gpu/host_offloading/gpu_host_offloading_allocator_test.cc +++ b/third_party/xla/xla/backends/gpu/host_offloading/gpu_host_offloading_allocator_test.cc @@ -31,7 +31,6 @@ limitations under the License. #include "xla/stream_executor/platform_manager.h" #include "xla/stream_executor/stream.h" #include "xla/stream_executor/stream_executor.h" -#include "xla/tsl/platform/statusor.h" namespace xla::gpu { @@ -47,18 +46,17 @@ se::StreamExecutor* GpuExecutor() { TEST(GpuHostOffloadingAllocatorTest, AllocateTransferBuffer) { se::StreamExecutor* stream_executor = GpuExecutor(); auto allocator = CreateGpuHostOffloadingAllocator(stream_executor); - TF_ASSERT_OK_AND_ASSIGN(auto buffer, allocator->AllocateTransferBuffer(1024)); + ASSERT_OK_AND_ASSIGN(auto buffer, allocator->AllocateTransferBuffer(1024)); EXPECT_EQ(buffer->size_bytes(), 1024); - TF_ASSERT_OK_AND_ASSIGN( - auto memory_type, - stream_executor->GetPointerMemorySpace(buffer->untyped_data())); + ASSERT_OK_AND_ASSIGN(auto memory_type, stream_executor->GetPointerMemorySpace( + buffer->untyped_data())); EXPECT_EQ(memory_type, stream_executor::MemorySpace::kHost); } TEST(GpuHostOffloadingAllocatorTest, AllocateStagingBuffer) { se::StreamExecutor* stream_executor = GpuExecutor(); auto allocator = CreateGpuHostOffloadingAllocator(stream_executor); - TF_ASSERT_OK_AND_ASSIGN(auto buffer, allocator->AllocateStagingBuffer(1024)); + ASSERT_OK_AND_ASSIGN(auto buffer, allocator->AllocateStagingBuffer(1024)); EXPECT_EQ(buffer->size_bytes(), 1024); auto memory_type_or_status = diff --git a/third_party/xla/xla/backends/gpu/libraries/cutedsl/BUILD b/third_party/xla/xla/backends/gpu/libraries/cutedsl/BUILD index c07f491334f1bc..66501ad24dbcdb 100644 --- a/third_party/xla/xla/backends/gpu/libraries/cutedsl/BUILD +++ b/third_party/xla/xla/backends/gpu/libraries/cutedsl/BUILD @@ -62,7 +62,6 @@ xla_test( ":ffi", "//xla:error_spec", "//xla/backends/gpu/tests:hlo_pjrt_gpu_test_base", - "//xla/tsl/lib/core:status_test_util", "//xla/tsl/platform:env", "//xla/tsl/platform:test", "@com_google_absl//absl/strings:string_view", diff --git a/third_party/xla/xla/backends/gpu/libraries/cutedsl/ffi_test.cc b/third_party/xla/xla/backends/gpu/libraries/cutedsl/ffi_test.cc index 4882d77d70db61..687d5fd72d7475 100644 --- a/third_party/xla/xla/backends/gpu/libraries/cutedsl/ffi_test.cc +++ b/third_party/xla/xla/backends/gpu/libraries/cutedsl/ffi_test.cc @@ -15,11 +15,11 @@ limitations under the License. #include +#include #include #include "absl/strings/string_view.h" #include "xla/backends/gpu/tests/hlo_pjrt_gpu_test_base.h" #include "xla/error_spec.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/tsl/platform/env.h" #include "xla/tsl/platform/test.h" #include "tsl/platform/path.h" @@ -34,7 +34,7 @@ TEST_F(CuteDslCustomCallTest, RunVectorAdd) { tsl::io::JoinPath(tsl::testing::XlaSrcRoot(), "backends", "gpu", "libraries", "cutedsl", "vector_add.hlo"); std::string hlo_text; - TF_ASSERT_OK(tsl::ReadFileToString(tsl::Env::Default(), hlo_path, &hlo_text)); + ASSERT_OK(tsl::ReadFileToString(tsl::Env::Default(), hlo_path, &hlo_text)); std::string reference_hlo_text = R"( HloModule reference, entry_computation_layout={(f32[1024]{0}, f32[1024]{0})->f32[1024]{0}} diff --git a/third_party/xla/xla/backends/gpu/profiler/BUILD b/third_party/xla/xla/backends/gpu/profiler/BUILD index 0a78b75eed60aa..4c3b24cc0b97b6 100644 --- a/third_party/xla/xla/backends/gpu/profiler/BUILD +++ b/third_party/xla/xla/backends/gpu/profiler/BUILD @@ -116,7 +116,6 @@ xla_test( "//xla/stream_executor/gpu:gpu_test_kernels", "//xla/stream_executor/gpu:gpu_test_kernels_fatbin", "//xla/stream_executor/rocm:rocm_platform_id", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_macros", "@com_google_absl//absl/status:status_matchers", diff --git a/third_party/xla/xla/backends/gpu/profiler/kernel_name_tracer_test.cc b/third_party/xla/xla/backends/gpu/profiler/kernel_name_tracer_test.cc index 914bbe532ff2c6..35da2bcfb122f7 100644 --- a/third_party/xla/xla/backends/gpu/profiler/kernel_name_tracer_test.cc +++ b/third_party/xla/xla/backends/gpu/profiler/kernel_name_tracer_test.cc @@ -55,7 +55,6 @@ limitations under the License. #include "xla/stream_executor/rocm/rocm_platform_id.h" #include "xla/stream_executor/stream.h" #include "xla/stream_executor/stream_executor_memory_allocator.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla_data.pb.h" namespace xla::gpu { @@ -75,10 +74,9 @@ absl::StatusOr GetPlatform() { class KernelNameTracerTest : public ::testing::Test { protected: void SetUp() override { - TF_ASSERT_OK_AND_ASSIGN(platform_, GetPlatform()); - TF_ASSERT_OK_AND_ASSIGN(stream_executor_, platform_->ExecutorForDevice(0)); - TF_ASSERT_OK_AND_ASSIGN(stream_, - stream_executor_->CreateStream(std::nullopt)); + ASSERT_OK_AND_ASSIGN(platform_, GetPlatform()); + ASSERT_OK_AND_ASSIGN(stream_executor_, platform_->ExecutorForDevice(0)); + ASSERT_OK_AND_ASSIGN(stream_, stream_executor_->CreateStream(std::nullopt)); } stream_executor::Platform* platform_; @@ -92,8 +90,8 @@ void LaunchAddI32Kernels(stream_executor::StreamExecutor* executor, stream_executor::TypedKernel, stream_executor::DeviceAddress, stream_executor::DeviceAddress>; - TF_ASSERT_OK_AND_ASSIGN(AddI32Kernel add, - stream_executor::gpu::LoadAddI32TestKernel(executor)); + ASSERT_OK_AND_ASSIGN(AddI32Kernel add, + stream_executor::gpu::LoadAddI32TestKernel(executor)); constexpr int64_t kLength = 4; constexpr int64_t kLengthInBytes = sizeof(int32_t) * kLength; @@ -131,8 +129,8 @@ void LaunchCommandBufferThunk(stream_executor::StreamExecutor* executor, stream_executor::TypedKernel, stream_executor::DeviceAddress, stream_executor::DeviceAddress>; - TF_ASSERT_OK_AND_ASSIGN(AddI32Kernel add, - stream_executor::gpu::LoadAddI32TestKernel(executor)); + ASSERT_OK_AND_ASSIGN(AddI32Kernel add, + stream_executor::gpu::LoadAddI32TestKernel(executor)); constexpr int64_t kLength = 4; constexpr int64_t kLengthInBytes = sizeof(int32_t) * kLength; @@ -171,11 +169,10 @@ void LaunchCommandBufferThunk(stream_executor::StreamExecutor* executor, commands.Append(KernelThunk::MakeKernelThunk("AddI32", args, args_access, LaunchDimensions(1, kLength), /*shmem_bytes=*/0)); - TF_ASSERT_OK_AND_ASSIGN( - CommandExecutor cmd_buffer_executor, - CommandExecutor::Create( - std::move(commands), - CommandExecutor::SynchronizationMode::kConcurrent)); + ASSERT_OK_AND_ASSIGN(CommandExecutor cmd_buffer_executor, + CommandExecutor::Create( + std::move(commands), + CommandExecutor::SynchronizationMode::kConcurrent)); // Construct a thunk with command sequence. CommandBufferThunk thunk(std::move(cmd_buffer_executor), Thunk::ThunkInfo()); @@ -190,9 +187,9 @@ void LaunchCommandBufferThunk(stream_executor::StreamExecutor* executor, /*persistent_alloc_indices=*/absl::Span()); // This is where we're getting the 'AddI32' kernel from. - TF_ASSERT_OK_AND_ASSIGN(std::vector fatbin, - stream_executor::gpu::GetGpuTestKernelsFatbin( - executor->GetPlatform()->Name())); + ASSERT_OK_AND_ASSIGN(std::vector fatbin, + stream_executor::gpu::GetGpuTestKernelsFatbin( + executor->GetPlatform()->Name())); ASSERT_THAT( thunk.Initialize({executor, Thunk::ExecutableSource{/*text=*/"", fatbin, @@ -210,7 +207,7 @@ void LaunchCommandBufferThunk(stream_executor::StreamExecutor* executor, } TEST_F(KernelNameTracerTest, Create) { - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( std::unique_ptr tracer, KernelNameTracer::Create(stream_executor::cuda::kCudaPlatformId)); tracer->start(); @@ -224,8 +221,8 @@ TEST_F(KernelNameTracerTest, CreateUnsupportedPlatform) { } TEST_F(KernelNameTracerTest, CaptureKernelNames) { - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr tracer, - KernelNameTracer::Create(platform_->id())); + ASSERT_OK_AND_ASSIGN(std::unique_ptr tracer, + KernelNameTracer::Create(platform_->id())); tracer->start(); LaunchAddI32Kernels(stream_executor_, stream_.get()); @@ -235,8 +232,8 @@ TEST_F(KernelNameTracerTest, CaptureKernelNames) { } TEST_F(KernelNameTracerTest, CaptureKernelNamesFromCommandBufferThunk) { - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr tracer, - KernelNameTracer::Create(platform_->id())); + ASSERT_OK_AND_ASSIGN(std::unique_ptr tracer, + KernelNameTracer::Create(platform_->id())); tracer->start(); LaunchCommandBufferThunk(stream_executor_, stream_.get()); diff --git a/third_party/xla/xla/backends/gpu/target_config/BUILD b/third_party/xla/xla/backends/gpu/target_config/BUILD index d44a0855fc5b96..b6f29bfd458d0c 100644 --- a/third_party/xla/xla/backends/gpu/target_config/BUILD +++ b/third_party/xla/xla/backends/gpu/target_config/BUILD @@ -82,7 +82,6 @@ xla_cc_test( ":target_config", "//xla/stream_executor:device_description_proto_cc", "//xla/stream_executor:semantic_version", - "//xla/tsl/lib/core:status_test_util", "//xla/tsl/platform:env", "//xla/tsl/platform:status_matchers", "@com_google_absl//absl/status", @@ -136,7 +135,6 @@ xla_test( "//xla/stream_executor:platform", "//xla/stream_executor:platform_manager", "//xla/stream_executor:stream_executor_h", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/log", "@com_google_absl//absl/strings", @@ -214,9 +212,7 @@ xla_test( "//xla/stream_executor:platform", "//xla/stream_executor:platform_manager", "//xla/stream_executor:stream_executor_h", - "//xla/tsl/lib/core:status_test_util", "//xla/tsl/platform:env", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/log", diff --git a/third_party/xla/xla/backends/gpu/target_config/cudnn_device_props_test.cc b/third_party/xla/xla/backends/gpu/target_config/cudnn_device_props_test.cc index 0182f966175ef4..b6d77656218b2f 100644 --- a/third_party/xla/xla/backends/gpu/target_config/cudnn_device_props_test.cc +++ b/third_party/xla/xla/backends/gpu/target_config/cudnn_device_props_test.cc @@ -21,6 +21,7 @@ limitations under the License. #include #include +#include #include #include "absl/log/log.h" #include "absl/strings/ascii.h" @@ -33,7 +34,6 @@ limitations under the License. #include "xla/stream_executor/platform.h" #include "xla/stream_executor/platform_manager.h" #include "xla/stream_executor/stream_executor.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" namespace xla::gpu { @@ -116,13 +116,13 @@ TEST(CudnnDevicePropsTest, MatchesLiveDevice) { std::string name = absl::AsciiStrToUpper(PlatformUtil::CanonicalPlatformName("gpu").value()); - TF_ASSERT_OK_AND_ASSIGN(se::Platform * platform, - se::PlatformManager::PlatformWithName(name)); + ASSERT_OK_AND_ASSIGN(se::Platform * platform, + se::PlatformManager::PlatformWithName(name)); bool any_compared = false; for (int i = 0; i < platform->VisibleDeviceCount(); ++i) { - TF_ASSERT_OK_AND_ASSIGN(se::StreamExecutor * executor, - platform->ExecutorForDevice(i)); + ASSERT_OK_AND_ASSIGN(se::StreamExecutor * executor, + platform->ExecutorForDevice(i)); const se::DeviceDescription& desc = executor->GetDeviceDescription(); SCOPED_TRACE(absl::StrCat("device ", i, ": ", desc.name())); @@ -130,7 +130,7 @@ TEST(CudnnDevicePropsTest, MatchesLiveDevice) { auto build_err = live->set_device_id(i).build(); ASSERT_FALSE(build_err.is_bad()) << build_err.get_message(); - TF_ASSERT_OK_AND_ASSIGN(auto synth, BuildDeviceProperties(desc)); + ASSERT_OK_AND_ASSIGN(auto synth, BuildDeviceProperties(desc)); Json::Value live_json = ParseProps(live); Json::Value synth_json = ParseProps(synth); diff --git a/third_party/xla/xla/backends/gpu/target_config/embedded_target_config_test.cc b/third_party/xla/xla/backends/gpu/target_config/embedded_target_config_test.cc index 8e74453e01872a..77bf1e1e19a4e8 100644 --- a/third_party/xla/xla/backends/gpu/target_config/embedded_target_config_test.cc +++ b/third_party/xla/xla/backends/gpu/target_config/embedded_target_config_test.cc @@ -15,6 +15,7 @@ limitations under the License. #include +#include #include #include "absl/container/flat_hash_map.h" #include "absl/log/log.h" @@ -26,9 +27,7 @@ limitations under the License. #include "xla/stream_executor/platform.h" #include "xla/stream_executor/platform_manager.h" #include "xla/stream_executor/stream_executor.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/tsl/platform/env.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "tsl/platform/path.h" #include "tsl/platform/platform.h" @@ -45,7 +44,7 @@ TEST(EmbeddedTargetConfigTest, DeviceInfoMatches) { "rtx6000pro", "gb200", "gb300"}) { GpuTargetConfigProto proto; std::string spec_string; - TF_ASSERT_OK(tsl::ReadFileToString( + ASSERT_OK(tsl::ReadFileToString( tsl::Env::Default(), tsl::io::JoinPath(tsl::testing::XlaSrcRoot(), "backends/gpu/target_config/specs", @@ -57,14 +56,14 @@ TEST(EmbeddedTargetConfigTest, DeviceInfoMatches) { } auto name = absl::AsciiStrToUpper( xla::PlatformUtil::CanonicalPlatformName("gpu").value()); - TF_ASSERT_OK_AND_ASSIGN(Platform * platform, - PlatformManager::PlatformWithName(name)); + ASSERT_OK_AND_ASSIGN(Platform * platform, + PlatformManager::PlatformWithName(name)); ASSERT_GT(platform->VisibleDeviceCount(), 0) << "No visible GPU devices found (cuInit may have failed)."; bool all_skipped = true; for (int i = 0; i < platform->VisibleDeviceCount(); ++i) { - TF_ASSERT_OK_AND_ASSIGN(StreamExecutor * executor, - platform->ExecutorForDevice(i)); + ASSERT_OK_AND_ASSIGN(StreamExecutor * executor, + platform->ExecutorForDevice(i)); const DeviceDescription& physical_device_description = executor->GetDeviceDescription(); diff --git a/third_party/xla/xla/backends/gpu/target_config/target_config_test.cc b/third_party/xla/xla/backends/gpu/target_config/target_config_test.cc index d81bad52a591b4..00acd2106f7f26 100644 --- a/third_party/xla/xla/backends/gpu/target_config/target_config_test.cc +++ b/third_party/xla/xla/backends/gpu/target_config/target_config_test.cc @@ -23,7 +23,6 @@ limitations under the License. #include "google/protobuf/text_format.h" #include "xla/stream_executor/device_description.pb.h" #include "xla/stream_executor/semantic_version.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/tsl/platform/env.h" #include "xla/tsl/platform/status_matchers.h" #include "tsl/platform/path.h" @@ -106,7 +105,7 @@ TEST(TargetConfigTest, GetTargetConfigFromFile) { platform_name: "platform" gpu_device_info { threads_per_block_limit: 5 } )pb"; - TF_ASSERT_OK( + ASSERT_OK( tsl::WriteStringToFile(tsl::Env::Default(), filename, proto_content)); ASSERT_OK_AND_ASSIGN(GpuTargetConfig config, diff --git a/third_party/xla/xla/backends/gpu/tests/BUILD b/third_party/xla/xla/backends/gpu/tests/BUILD index c439f9a6457636..6bfdb98d77f5ba 100644 --- a/third_party/xla/xla/backends/gpu/tests/BUILD +++ b/third_party/xla/xla/backends/gpu/tests/BUILD @@ -165,7 +165,6 @@ xla_test( "//xla/tests:hlo_pjrt_interpreter_reference_mixin", "//xla/tests:hlo_pjrt_test_base", "//xla/tests:literal_test_util", - "//xla/tsl/platform:statusor", "@com_google_googletest//:gtest_main", ], ) @@ -223,7 +222,6 @@ xla_test( "//xla:xla_proto_cc", "//xla/service:hlo_module_config", "//xla/tests:literal_test_util", - "//xla/tsl/platform:statusor", "@com_google_googletest//:gtest_main", ], ) @@ -261,7 +259,6 @@ xla_test( "//xla/service:hlo_module_config", "//xla/service:hlo_runner_interface", "//xla/tests:hlo_pjrt_test_base", - "//xla/tsl/lib/core:status_test_util", "@com_google_absl//absl/algorithm:container", "@com_google_absl//absl/status:statusor", "@com_google_googletest//:gtest_main", @@ -433,7 +430,6 @@ xla_test( "//xla/stream_executor:device_description", "//xla/stream_executor/cuda:cuda_compute_capability", "//xla/tests:literal_test_util", - "//xla/tsl/lib/core:status_test_util", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/status:statusor", @@ -612,7 +608,6 @@ xla_test( deps = [ ":gpu_pjrt_codegen_test", "//xla/stream_executor/cuda:cuda_compute_capability", - "//xla/tsl/lib/core:status_test_util", "@com_google_googletest//:gtest_main", ], ) @@ -662,7 +657,6 @@ xla_test( "//xla/tests:hlo_pjrt_interpreter_reference_mixin", "//xla/tests:hlo_pjrt_test_base", "//xla/tests:test_utils", - "//xla/tsl/platform:statusor", "@com_google_googletest//:gtest_main", ], ) @@ -719,8 +713,8 @@ xla_test( "//xla/hlo/ir:hlo", "//xla/tests:hlo_pjrt_interpreter_reference_mixin", "//xla/tests:hlo_pjrt_test_base", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/log:check", + "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", "@com_google_googletest//:gtest_main", "@eigen_archive//:eigen3", @@ -1021,10 +1015,9 @@ cc_library( deps = [ "//xla/hlo/testlib:verified_hlo_module", "//xla/tests:hlo_pjrt_test_base", - "//xla/tsl/lib/core:status_test_util", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/strings", + "@com_google_googletest//:gtest_for_library", ], alwayslink = True, # This library registers test cases at static initialization time. ) @@ -1125,7 +1118,6 @@ xla_test( "//xla:literal_util", "//xla/tests:hlo_pjrt_test_base", "//xla/tests:literal_test_util", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_googletest//:gtest_main", ], @@ -1139,7 +1131,6 @@ xla_test( deps = [ "//xla:literal", "//xla/tests:hlo_pjrt_test_base", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/strings:string_view", "@com_google_googletest//:gtest_main", ], @@ -1248,7 +1239,6 @@ xla_test( "//xla/tests:test_utils", "//xla/tests:xla_internal_test_main", "//xla/tests/restricted:hlo_test_base_legacy", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/log", "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/types:span", @@ -1275,10 +1265,10 @@ xla_test( "//xla/tests:hlo_pjrt_test_base", "//xla/tests:xla_internal_test_main", "//xla/tsl/platform:logging", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/types:span", + "@com_google_googletest//:gtest", ], ) @@ -1320,9 +1310,7 @@ xla_test( "//xla/tests:literal_test_util", "//xla/tests:test_utils", "//xla/tests:xla_internal_test_main", - "//xla/tsl/lib/core:status_test_util", "//xla/tsl/platform:status_matchers", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/base:nullability", "@com_google_absl//absl/log", @@ -1332,6 +1320,7 @@ xla_test( "@com_google_absl//absl/strings", "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/types:span", + "@com_google_googletest//:gtest", ], ) @@ -1475,6 +1464,7 @@ xla_test( "@com_google_absl//absl/synchronization", "@com_google_absl//absl/time", "@com_google_absl//absl/types:span", + "@com_google_googletest//:gtest", ] + if_cuda_is_configured([ ":collective_ops_ffi_kernels_cuda", "@com_google_absl//absl/base", @@ -1507,7 +1497,6 @@ xla_test( "//xla/tests:literal_test_util", "//xla/tests:test_utils", "//xla/tests:xla_internal_test_main", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/algorithm:container", "@com_google_absl//absl/log", @@ -1516,6 +1505,7 @@ xla_test( "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/types:span", + "@com_google_googletest//:gtest", "@tsl//tsl/platform:regexp", ], ) @@ -1595,11 +1585,10 @@ xla_test( "//xla/tests:hlo_pjrt_test_base", "//xla/tests:literal_test_util", "//xla/tests:xla_internal_test_main", - "//xla/tsl/lib/core:status_test_util", "//xla/tsl/platform:logging", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/strings:string_view", + "@com_google_googletest//:gtest", ], ) @@ -1638,11 +1627,11 @@ xla_test( "//xla/tests:pjrt_client_registry", "//xla/tests:xla_internal_test_main", # fixdeps: keep "//xla/tests:xla_test_backend_predicates", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/log:check", "@com_google_absl//absl/strings", "@com_google_absl//absl/types:span", + "@com_google_googletest//:gtest", "@com_google_protobuf//:protobuf", ], ) @@ -1678,7 +1667,6 @@ xla_test( "//xla/service/gpu:backend_configs_cc", "//xla/stream_executor/gpu:all_reduce_kernel", "//xla/tests:literal_test_util", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "//xla/tsl/testing:temporary_directory", "@com_google_absl//absl/log", @@ -1911,8 +1899,8 @@ xla_test( "//xla/service:hlo_module_config", "//xla/tests:literal_test_util", "//xla/tests:xla_internal_test_main", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/strings:string_view", + "@com_google_googletest//:gtest", ], ) diff --git a/third_party/xla/xla/backends/gpu/tests/all_reduce_e2e_test.cc b/third_party/xla/xla/backends/gpu/tests/all_reduce_e2e_test.cc index 72f210749a19ce..051f757fb068c2 100644 --- a/third_party/xla/xla/backends/gpu/tests/all_reduce_e2e_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/all_reduce_e2e_test.cc @@ -56,7 +56,6 @@ limitations under the License. #include "xla/shape_util.h" #include "xla/stream_executor/gpu/all_reduce_kernel.h" #include "xla/tests/literal_test_util.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "xla/tsl/testing/temporary_directory.h" #include "xla/types.h" @@ -613,16 +612,16 @@ TEST_P(AllReduceTypesTest, SupportedTypes2GPUs) { kModuleStr, primitive_util::LowercasePrimitiveTypeName(element_type), absl::StrJoin(shape, ","), HloOpcodeString(opcode)); SCOPED_TRACE(::testing::Message() << "module_str: " << module_str); - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(module_str, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(module_str, kNumReplicas)); + ASSERT_OK_AND_ASSIGN( InputsOutputs test_io, (BuildTestInputsOutputs(element_type, opcode, *module, kNumReplicas, /*num_iterations=*/1))); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result, ExecuteReplicated(std::move(module), - /*arguments=*/test_io.InputLiteralPtrs())) + /*arguments=*/test_io.InputLiteralPtrs())); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); for (int i = 0; i < kNumReplicas; ++i) { @@ -656,13 +655,13 @@ TEST_P(AllReduceLayoutAwareTest, AllReduceLayoutAwareTest) { kModuleStr, primitive_util::LowercasePrimitiveTypeName(element_type), absl::StrJoin(shape, ","), layout_str, HloOpcodeString(opcode)); SCOPED_TRACE(::testing::Message() << "module_str: " << module_str); - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(module_str, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(module_str, kNumReplicas)); + ASSERT_OK_AND_ASSIGN( InputsOutputs test_io, (BuildTestInputsOutputs(element_type, opcode, *module, kNumReplicas, /*num_iterations=*/1))); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result, ExecuteReplicated(std::move(module), /*arguments=*/test_io.InputLiteralPtrs())); @@ -697,19 +696,19 @@ TEST_P(AllReduceTest, F32_8GPUs_AllReplicasOneGroup) { } const std::vector shape = GetParam().GetShape(PrimitiveType::F32); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule( absl::StrFormat(kModuleStr, absl::StrJoin(shape, ",")), kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( InputsOutputs test_io, (BuildTestInputsOutputs( HloOpcode::kAdd, *module, kNumReplicas, /*num_iterations=*/1))); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result, ExecuteReplicated(std::move(module), - /*arguments=*/test_io.InputLiteralPtrs())) + /*arguments=*/test_io.InputLiteralPtrs())); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); for (int i = 0; i < kNumReplicas; ++i) { @@ -766,15 +765,15 @@ TEST_P(AllReduceTest, F32_8GPUs_2ReplicasPerGroup) { return; } - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( InputsOutputs test_io, (BuildTestInputsOutputs( HloOpcode::kAdd, *module, kNumReplicas, kNumIterations))); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result, ExecuteReplicated(std::move(module), /*arguments=*/test_io.InputLiteralPtrs())); @@ -811,19 +810,19 @@ TEST_P(AllReduceTest, F32TwoD4GPUs) { const std::vector shape = GetParam().GetShape(PrimitiveType::F32, /*rank=*/2); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule( absl::StrFormat(kModuleStr, absl::StrJoin(shape, ",")), kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( InputsOutputs test_io, (BuildTestInputsOutputs( HloOpcode::kAdd, *module, kNumReplicas, /*num_iterations=*/1))); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result, ExecuteReplicated(std::move(module), - /*arguments=*/test_io.InputLiteralPtrs())) + /*arguments=*/test_io.InputLiteralPtrs())); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); for (int i = 0; i < kNumReplicas; ++i) { @@ -855,19 +854,19 @@ TEST_P(AllReduceTest, F32_3D_2GPUs) { } const std::vector shape = GetParam().GetShape(PrimitiveType::F32, /*rank=*/3); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule( absl::StrFormat(kModuleStr, absl::StrJoin(shape, ",")), kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( InputsOutputs test_io, (BuildTestInputsOutputs( HloOpcode::kAdd, *module, kNumReplicas, /*num_iterations=*/1))); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result, ExecuteReplicated(std::move(module), - /*arguments=*/test_io.InputLiteralPtrs())) + /*arguments=*/test_io.InputLiteralPtrs())); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); for (int i = 0; i < kNumReplicas; ++i) { @@ -959,9 +958,9 @@ TEST_F(AllReduceTestNoParams, AsyncAllReduce_F8E4M3FN_TrainingStep_2GPUs) { Literal upstream_grad_lit1 = LiteralUtil::CreateFromArray(upstream_grad1); Literal upstream_grad_lit2 = LiteralUtil::CreateFromArray(upstream_grad2); - TF_ASSERT_OK_AND_ASSIGN(auto f16_module, ParseAndReturnVerifiedModule( - kF16ModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(auto f16_module, ParseAndReturnVerifiedModule( + kF16ModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN( ExecutionResult f16_result, ExecuteReplicated( std::move(f16_module), @@ -971,9 +970,9 @@ TEST_F(AllReduceTestNoParams, AsyncAllReduce_F8E4M3FN_TrainingStep_2GPUs) { // Verify FP16 all-reduce type in optimized module VerifyAllReduceType(f16_result.optimized_module, F16); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto f8_module, ParseAndReturnVerifiedModule(kF8ModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult f8_result, ExecuteReplicated( std::move(f8_module), @@ -1000,8 +999,8 @@ TEST_F(AllReduceTestNoParams, AsyncAllReduce_F8E4M3FN_TrainingStep_2GPUs) { // Numerical precision check: FP8 should produce measurably different // results than FP16. FP8 e4m3 has ~6% relative error (2^-4), FP16 has ~0.1% // (2^-10). - TF_ASSERT_OK_AND_ASSIGN(Literal f16_f32, f16_r0[1].Convert(F32)); - TF_ASSERT_OK_AND_ASSIGN(Literal f8_f32, f8_r0[1].Convert(F32)); + ASSERT_OK_AND_ASSIGN(Literal f16_f32, f16_r0[1].Convert(F32)); + ASSERT_OK_AND_ASSIGN(Literal f8_f32, f8_r0[1].Convert(F32)); absl::Span f16_data = f16_f32.data(); absl::Span f8_data = f8_f32.data(); float max_abs_diff = 0.0f; @@ -1043,7 +1042,7 @@ TEST_F(AllReduceTestNoParams, AsyncAllReduce_F8E4M3FN_FailsOnUnsupportedGPUs) { return; } - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(kF8ModuleStr, kNumReplicas)); Array input1({64, 128}), input2({64, 128}); diff --git a/third_party/xla/xla/backends/gpu/tests/async_command_buffer_test.cc b/third_party/xla/xla/backends/gpu/tests/async_command_buffer_test.cc index e86c3ae9133481..64988ced79aebe 100644 --- a/third_party/xla/xla/backends/gpu/tests/async_command_buffer_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/async_command_buffer_test.cc @@ -15,6 +15,7 @@ limitations under the License. #include +#include #include #include "xla/backends/gpu/tests/hlo_pjrt_gpu_test_base.h" #include "xla/debug_options_flags.h" @@ -22,7 +23,6 @@ limitations under the License. #include "xla/literal_util.h" #include "xla/service/hlo_module_config.h" #include "xla/tests/literal_test_util.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla.pb.h" namespace xla::gpu { @@ -85,9 +85,8 @@ TEST_F(AsyncCommandBufferTest, CommandBuffer) { Literal argument = LiteralUtil::CreateR2({{1.0, 2.0}, {3.0, 4.0}}); Literal expected = LiteralUtil::CreateR2({{4.0, 8.0}, {12.0, 16.0}}); - TF_ASSERT_OK_AND_ASSIGN( - Literal result, - Execute(std::move(module), {&argument}, /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {&argument}, + /*run_hlo_passes=*/false)); EXPECT_TRUE(LiteralTestUtil::Equal(expected, result)); } diff --git a/third_party/xla/xla/backends/gpu/tests/async_kernel_launch_test.cc b/third_party/xla/xla/backends/gpu/tests/async_kernel_launch_test.cc index abadce185cd162..1b872ba8671de1 100644 --- a/third_party/xla/xla/backends/gpu/tests/async_kernel_launch_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/async_kernel_launch_test.cc @@ -15,6 +15,7 @@ limitations under the License. #include +#include #include #include "xla/debug_options_flags.h" #include "xla/error_spec.h" @@ -24,7 +25,6 @@ limitations under the License. #include "xla/tests/hlo_pjrt_interpreter_reference_mixin.h" #include "xla/tests/hlo_pjrt_test_base.h" #include "xla/tests/literal_test_util.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla.pb.h" namespace xla::gpu { @@ -76,9 +76,8 @@ TEST_F(AsyncKernelLaunchTest, BasicFusion) { Literal argument = LiteralUtil::CreateR2({{1.0, 2.0}, {3.0, 4.0}}); Literal expected = LiteralUtil::CreateR2({{4.0, 8.0}, {12.0, 16.0}}); - TF_ASSERT_OK_AND_ASSIGN( - Literal result, - Execute(std::move(module), {&argument}, /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {&argument}, + /*run_hlo_passes=*/false)); EXPECT_TRUE(LiteralTestUtil::Equal(expected, result)); } diff --git a/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test.cc b/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test.cc index 0ad0026d5c032a..cd1b6af993ad11 100644 --- a/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/collective_ops_e2e_test.cc @@ -23,6 +23,7 @@ limitations under the License. #include #include +#include #include "absl/base/nullability.h" #include "absl/log/check.h" #include "absl/log/log.h" @@ -60,9 +61,7 @@ limitations under the License. #include "xla/stream_executor/device_description.h" #include "xla/tests/literal_test_util.h" #include "xla/tests/test_utils.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/tsl/platform/status_matchers.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "xla/types.h" #include "xla/xla.pb.h" @@ -163,15 +162,14 @@ class CollectiveOpsTestE2E : public CollectiveOpsE2ETestBase { HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); config.set_debug_options(options); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(hlo_text, config)); - - TF_ASSERT_OK_AND_ASSIGN(auto executable, - CreateExecutable(std::move(module), - /*run_hlo_passes=*/true)); - TF_ASSERT_OK_AND_ASSIGN( - const HloModule* const hlo_module, - test_runner().HloModuleFromWrapped(executable.get())); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_text, config)); + + ASSERT_OK_AND_ASSIGN(auto executable, + CreateExecutable(std::move(module), + /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(const HloModule* const hlo_module, + test_runner().HloModuleFromWrapped(executable.get())); std::vector gemm_ops = FindInstructions(hlo_module, HloOpcode::kCustomCall); for (HloInstruction* gemm_op : gemm_ops) { @@ -358,11 +356,11 @@ TEST_P(AsyncCollectiveOps, AsyncAllReduce) { const bool enable_async_all_reduce = enable_async_; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async_all_reduce) { @@ -477,11 +475,11 @@ TEST_P(AsyncCollectiveOps, AsyncCollectiveBroadcast) { << device_count() << " available)"; const bool enable_async_collective_broadcast = enable_async_; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async_collective_broadcast) { @@ -529,13 +527,13 @@ TEST_P(AsyncCollectiveOps, AsyncCollectiveBroadcastDynamicRoot) { // (root rank -> broadcast value seen by every replica). for (const auto& [root_rank, expected] : std::vector>{{0, 10}, {1, 11}}) { - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule( absl::Substitute(kModuleTemplate, root_rank), kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async_) { @@ -588,11 +586,11 @@ TEST_P(AsyncCollectiveOps, AsyncCollectiveReduce) { << "Test requires at least " << kNumReplicas << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async_) { @@ -640,11 +638,11 @@ TEST_P(AsyncCollectiveOps, AsyncCollectiveReduceDynamicRoot) { << "Test requires at least " << kNumReplicas << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async_) { @@ -689,11 +687,11 @@ TEST_P(CollectivesModeOps, AllGather) { << "Test requires at least " << kNumReplicas << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async()) { @@ -736,11 +734,11 @@ TEST_P(CollectivesModeOps, AllGatherMixedTypes) { << "Test requires at least " << kNumReplicas << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async()) { @@ -884,11 +882,11 @@ TEST_P(CollectivesModeOps, CollectivePermute) { << "Test requires at least " << kNumReplicas << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async()) { @@ -926,16 +924,16 @@ TEST_P(CollectivesModeOps, CollectivePermuteOnParameters) { << "Test requires at least " << kNumReplicas << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); // Replica 0 gets {10, 10}, replica 1 gets {11, 11}. auto arg0 = LiteralUtil::CreateR1({10, 10}); auto arg1 = LiteralUtil::CreateR1({11, 11}); std::vector> args = {{&arg0}, {&arg1}}; - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module), args)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), args)); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); @@ -961,11 +959,11 @@ TEST_P(CollectivesModeOps, CombinedCollectivePermute) { )"; const int64_t kNumReplicas = 2; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async()) { @@ -1008,10 +1006,10 @@ TEST_P(CollectivesModeOps, CollectivePermuteCombiner) { << device_count() << " available)"; } - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async()) { @@ -1081,8 +1079,8 @@ TEST_F(CollectiveOpsTestE2E, CollectiveGroupAllReduceDifferentReplicaGroups) { HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); std::vector all_args; std::vector pair_args; @@ -1102,8 +1100,8 @@ TEST_F(CollectiveOpsTestE2E, CollectiveGroupAllReduceDifferentReplicaGroups) { args[replica] = {&all_args[replica], &pair_args[replica]}; } - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module), args)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), args)); std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); @@ -1300,11 +1298,11 @@ TEST_P(AsyncCollectiveOps, AsyncReduceScatter) { << device_count() << " available)"; const bool enable_async_reduce_scatter = enable_async_; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async_reduce_scatter) { @@ -1345,11 +1343,11 @@ TEST_P(AsyncCollectiveOps, AsyncAllToAllWithSplitDim) { << device_count() << " available)"; const bool enable_async_all_to_all = enable_async_; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async_all_to_all) { @@ -1389,10 +1387,10 @@ TEST_F(CollectiveOpsTestE2E, AsyncAllToAllMemCpyWithSplitDim) { GetModuleConfigForTest(/*replica_count=*/kNumReplicas); config.mutable_debug_options().set_xla_gpu_use_memcpy_local_p2p(true); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* executable_module = execution_result.optimized_module; // Verify that the all-to-all is not decomposed into a tuple all-to-all. @@ -1433,11 +1431,11 @@ TEST_P(AsyncCollectiveOps, AsyncAllToAllWithoutSplitDim) { << device_count() << " available)"; const bool enable_async_all_to_all = enable_async_; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async_all_to_all) { @@ -1485,11 +1483,11 @@ TEST_P(AsyncCollectiveOps, AsyncAllToAllMemCpyWithoutSplitDim) { GetModuleConfigForTest(/*replica_count=*/kNumReplicas); config.mutable_debug_options().set_xla_gpu_use_memcpy_local_p2p(true); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); LiteralTestUtil::ExpectR1Equal({10, 15, 11, 16}, results[0]); @@ -1513,11 +1511,11 @@ TEST_P(AsyncCollectiveOps, AsyncAllToAllNumberOfElementsLargerThanInt32Max) { << device_count() << " available)"; const bool enable_async_all_to_all = enable_async_; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; if (enable_async_all_to_all) { @@ -1574,11 +1572,11 @@ ENTRY entry { << "Test requires at least " << kNumReplicas << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; const bool enable_async_ragged_all_to_all = enable_async_; @@ -1636,11 +1634,11 @@ TEST_P(AsyncMemcpyCollectiveOps, AsyncAllToAllMultipleReplicaGroups) { HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); LiteralTestUtil::ExpectR1Equal({10, 13}, results[0]); @@ -1668,11 +1666,11 @@ TEST_P(AsyncMemcpyCollectiveOps, AsyncAllToAllDegenerateWithSplitDim) { HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); LiteralTestUtil::ExpectR1Equal({10, 20}, results[0]); @@ -1699,11 +1697,11 @@ TEST_P(AsyncMemcpyCollectiveOps, AsyncAllToAllDegenerateWithoutSplitDim) { HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); LiteralTestUtil::ExpectR1Equal({10, 20}, results[0]); @@ -1733,11 +1731,11 @@ TEST_P(MemcpyCollectiveOps, AllToAll8Gpus) { HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const std::vector& results = execution_result.results; Array expected({16}); @@ -1875,11 +1873,11 @@ TEST_P(CollectivesModeOps, CollectivePermuteInWhileLoop) { ASSERT_GE(device_count(), kNumReplicas) << "Test requires at least " << kNumReplicas << " devices"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); // Trace through 4 iterations with replica 0 (r=0) and replica 1 (r=1): // Init: r0=[10,10] r1=[11,11] @@ -1942,11 +1940,11 @@ TEST_P(CollectivesModeOps, CombinedCollectivePermuteInWhileLoop) { ASSERT_GE(device_count(), kNumReplicas) << "Test requires at least " << kNumReplicas << " devices"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); // 4 iterations of pure swap (even number → back to original): // Init: r0=[10,10],[10.0,10.0] r1=[11,11],[11.0,11.0] @@ -2064,10 +2062,10 @@ TEST_F(CollectiveOpsTestE2E, WhileLoopReduceScatterCodeMotion) { config.mutable_debug_options() .set_xla_gpu_enable_while_loop_reduce_scatter_code_motion(true); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* executable_module = execution_result.optimized_module; @@ -2110,11 +2108,11 @@ TEST_F(CollectiveOpsTestE2E, NoAllToAllDecomposition) { HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* executable_module = execution_result.optimized_module; // Verify that the all-to-all is not decomposed into a tuple all-to-all. @@ -2158,14 +2156,14 @@ TEST_F(CollectiveOpsTestE2E, NoAsyncCollectives) { "gpu-convert-async-collectives-to-sync"); config.mutable_debug_options().add_xla_gpu_disable_async_collectives( DebugOptions::ALLCOLLECTIVES); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(auto executable, - CreateExecutable(std::move(module), - /*run_hlo_passes=*/true)); - TF_ASSERT_OK_AND_ASSIGN(const HloModule* const executable_module, - test_runner().HloModuleFromWrapped(executable.get())); + ASSERT_OK_AND_ASSIGN(auto executable, + CreateExecutable(std::move(module), + /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(const HloModule* const executable_module, + test_runner().HloModuleFromWrapped(executable.get())); // Verify that the all-to-all is a sync collective. const HloInstruction* all_to_all = @@ -2196,10 +2194,10 @@ TEST_F(CollectiveOpsTestE2E, HostMemoryOffloadingWithDonation) { config.mutable_debug_options().set_xla_gpu_enable_host_memory_offloading( true); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnUnverifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnUnverifiedModule(kModuleStr, config)); - TF_ASSERT_OK(module->input_output_alias_config().SetUpAlias( + ASSERT_OK(module->input_output_alias_config().SetUpAlias( /*output_index=*/{}, /*param_number=*/0, /*param_index=*/{}, @@ -2246,8 +2244,8 @@ class CollectiveOpsTestE2EWindowedNonWindowed : public CollectiveOpsTestE2E { } // Run with reference config. - TF_ASSERT_OK_AND_ASSIGN(auto ref_module, - ParseAndReturnVerifiedModule(hlo_text, config)); + ASSERT_OK_AND_ASSIGN(auto ref_module, + ParseAndReturnVerifiedModule(hlo_text, config)); ASSERT_OK_AND_ASSIGN(auto ref_executable, CreateExecutable(std::move(ref_module), /*run_hlo_passes=*/true)); @@ -2272,12 +2270,11 @@ class CollectiveOpsTestE2EWindowedNonWindowed : public CollectiveOpsTestE2E { debug_options.set_xla_gpu_multi_streamed_windowed_einsum(true); debug_options.set_xla_gpu_experimental_enable_alltoall_windowed_einsum( enable_a2a_rewrite); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(hlo_text, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_text, config)); - TF_ASSERT_OK_AND_ASSIGN( - ExecutionResult execution_result, - ExecuteReplicated(std::move(module), ref_fake_ptrs)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), ref_fake_ptrs)); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumPartitions); @@ -2321,11 +2318,11 @@ TEST_F(CollectiveOpsTestE2E, CollectiveMultiStreaming) { HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); config.set_debug_options(debug_options); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* executable_module = execution_result.optimized_module; ASSERT_NE(executable_module, nullptr); ASSERT_TRUE(executable_module->has_schedule()); @@ -2913,16 +2910,16 @@ class CollectiveOpsTestE2EPipelinedNonPipelined : public CollectiveOpsTestE2E { HloModuleConfig config = GetModuleConfigForTest(kNumReplicas, kNumPartitions); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(hlo_string, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, config)); auto fake_arguments = xla::MakeFakeArguments(module.get()).value(); std::vector fake_ptrs(fake_arguments.size()); for (int i = 0; i < fake_arguments.size(); ++i) { fake_ptrs[i] = &fake_arguments[i]; } - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module), fake_ptrs)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), fake_ptrs)); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumPartitions); @@ -2936,15 +2933,15 @@ class CollectiveOpsTestE2EPipelinedNonPipelined : public CollectiveOpsTestE2E { ref_opts.set_xla_gpu_pipeline_reduce_scatter( DebugOptions::COLLECTIVE_PIPELINING_MODE_OFF); - TF_ASSERT_OK_AND_ASSIGN( - auto ref_module, ParseAndReturnVerifiedModule(hlo_string, ref_config)); + ASSERT_OK_AND_ASSIGN(auto ref_module, + ParseAndReturnVerifiedModule(hlo_string, ref_config)); auto fake_ref_arguments = xla::MakeFakeArguments(ref_module.get()).value(); std::vector ref_fake_ptrs(fake_ref_arguments.size()); for (int i = 0; i < fake_ref_arguments.size(); ++i) { ref_fake_ptrs[i] = &fake_ref_arguments[i]; } - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult ref_execution_result, ExecuteReplicated(std::move(ref_module), ref_fake_ptrs)); const std::vector& ref_results = ref_execution_result.results; @@ -3148,11 +3145,11 @@ ENTRY entry { << "Test requires at least " << kNumReplicas * kNumPartitions << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(kModuleReplicatedStr, kNumReplicas, kNumPartitions)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const HloModule* hlo_module = execution_result.optimized_module; HloInstruction* all_to_all = @@ -3187,12 +3184,12 @@ ENTRY entry { << "Test requires at least " << kNumReplicas * kNumPartitions << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(kModuleReplicatedStr, kNumReplicas, kNumPartitions)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); // Verify that the element type of the all-to-all has been changed to BF16. const HloModule* hlo_module = execution_result.optimized_module; @@ -3285,15 +3282,15 @@ ENTRY entry { HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(kModuleReplicatedStr, config)); - TF_ASSERT_OK_AND_ASSIGN(auto executable, - CreateExecutable(std::move(module), - /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(auto executable, + CreateExecutable(std::move(module), + /*run_hlo_passes=*/true)); - TF_ASSERT_OK_AND_ASSIGN(const HloModule* const hlo_module, - test_runner().HloModuleFromWrapped(executable.get())); + ASSERT_OK_AND_ASSIGN(const HloModule* const hlo_module, + test_runner().HloModuleFromWrapped(executable.get())); EXPECT_NE(hlo_module, nullptr); } @@ -3400,11 +3397,11 @@ ENTRY main.49 { GetModuleConfigForTest(kNumReplicas, kNumPartitions); ref_config.mutable_debug_options().set_xla_gpu_use_memcpy_local_p2p(false); - TF_ASSERT_OK_AND_ASSIGN(auto ref_module, - ParseAndReturnVerifiedModule(hlo_string, ref_config)); + ASSERT_OK_AND_ASSIGN(auto ref_module, + ParseAndReturnVerifiedModule(hlo_string, ref_config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult ref_execution_result, - ExecuteReplicated(std::move(ref_module), fake_ptrs)); + ASSERT_OK_AND_ASSIGN(ExecutionResult ref_execution_result, + ExecuteReplicated(std::move(ref_module), fake_ptrs)); const std::vector& ref_results = ref_execution_result.results; ASSERT_EQ(ref_results.size(), kNumPartitions); ErrorSpec error_spec{1e-5, 1e-5}; @@ -3453,8 +3450,8 @@ ENTRY main { "gpu-convert-async-collectives-to-sync"); config.set_use_spmd_partitioning(false); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(hlo_string, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, config)); auto fake_arguments = xla::MakeFakeArguments(module.get()).value(); std::vector fake_ptrs(fake_arguments.size()); for (int i = 0; i < fake_arguments.size(); ++i) { @@ -3465,22 +3462,21 @@ ENTRY main { GetModuleConfigForTest(kNumReplicas, kNumPartitions); ref_config.mutable_debug_options().set_xla_gpu_use_memcpy_local_p2p(false); - TF_ASSERT_OK_AND_ASSIGN(auto ref_module, - ParseAndReturnVerifiedModule(hlo_string, ref_config)); + ASSERT_OK_AND_ASSIGN(auto ref_module, + ParseAndReturnVerifiedModule(hlo_string, ref_config)); auto fake_ref_arguments = xla::MakeFakeArguments(ref_module.get()).value(); std::vector ref_fake_ptrs(fake_ref_arguments.size()); for (int i = 0; i < fake_ref_arguments.size(); ++i) { ref_fake_ptrs[i] = &fake_ref_arguments[i]; } - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module), fake_ptrs)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), fake_ptrs)); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumPartitions); - TF_ASSERT_OK_AND_ASSIGN( - ExecutionResult ref_execution_result, - ExecuteReplicated(std::move(ref_module), ref_fake_ptrs)); + ASSERT_OK_AND_ASSIGN(ExecutionResult ref_execution_result, + ExecuteReplicated(std::move(ref_module), ref_fake_ptrs)); const std::vector& ref_results = ref_execution_result.results; ASSERT_EQ(ref_results.size(), kNumPartitions); ErrorSpec error_spec{1e-5, 1e-5}; @@ -3520,13 +3516,13 @@ ENTRY main { "gpu-convert-async-collectives-to-sync"); config.set_use_spmd_partitioning(false); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(hlo_string, config)); - TF_ASSERT_OK_AND_ASSIGN(auto executable, - CreateExecutable(std::move(module), - /*run_hlo_passes=*/false)); - TF_ASSERT_OK_AND_ASSIGN(const HloModule* const executable_module, - test_runner().HloModuleFromWrapped(executable.get())); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, config)); + ASSERT_OK_AND_ASSIGN(auto executable, + CreateExecutable(std::move(module), + /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(const HloModule* const executable_module, + test_runner().HloModuleFromWrapped(executable.get())); const HloInstruction* ag_start = FindCollectiveStarts(executable_module, HloOpcode::kAllGather).at(0); // Both ag and its producer should have collective memory space. @@ -3569,13 +3565,13 @@ ROOT tuple = (bf16[1024,1024]{1,0}, bf16[]) tuple(all-reduce-done, all-reduce-do "gpu-convert-async-collectives-to-sync"); config.set_use_spmd_partitioning(false); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(hlo_string, config)); - TF_ASSERT_OK_AND_ASSIGN(auto executable, - CreateExecutable(std::move(module), - /*run_hlo_passes=*/false)); - TF_ASSERT_OK_AND_ASSIGN(const HloModule* const executable_module, - test_runner().HloModuleFromWrapped(executable.get())); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, config)); + ASSERT_OK_AND_ASSIGN(auto executable, + CreateExecutable(std::move(module), + /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(const HloModule* const executable_module, + test_runner().HloModuleFromWrapped(executable.get())); std::vector all_ar = FindCollectiveStarts(executable_module, HloOpcode::kAllReduce); // Both allreduces should have their operands copied to collective memory @@ -3593,18 +3589,18 @@ TEST_F(CollectiveOpsTestE2E, OptimizedSubByteAllGatherOnDim0OutputIsCorrect) { << "Test requires at least " << kNumReplicas << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN(auto unoptimized_module, - ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(auto unoptimized_module, + ParseAndReturnVerifiedModule(R"( HloModule m, replica_count=2 e { a = s4[2,4]{1,0:E(4)} constant({{0,1,2,3},{4,5,5,4}}) b = s4[4,4]{1,0:E(4)} all-gather(a), dimensions={0} })", - kNumReplicas)); + kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(unoptimized_module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(unoptimized_module))); const HloModule* module = execution_result.optimized_module; EXPECT_THAT(module->entry_computation()->root_instruction(), @@ -3630,18 +3626,18 @@ TEST_F(CollectiveOpsTestE2E, OptimizedSubByteAllGatherOnDim1OutputIsCorrect) { << "Test requires at least " << kNumReplicas << " devices (" << device_count() << " available)"; - TF_ASSERT_OK_AND_ASSIGN(auto unoptimized_module, - ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(auto unoptimized_module, + ParseAndReturnVerifiedModule(R"( HloModule m, replica_count=2 e { a = s4[4,2]{1,0:E(4)} constant({{0,1},{2,3},{4,5},{5,4}}) b = s4[4,4]{1,0:E(4)} all-gather(a), dimensions={1} })", - kNumReplicas)); + kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(unoptimized_module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(unoptimized_module))); const HloModule* module = execution_result.optimized_module; const HloInstruction* root = module->entry_computation()->root_instruction(); @@ -3670,19 +3666,19 @@ TEST_F(CollectiveOpsTestE2E, AllGatherOnChangedDimensionIsCorrect) { ASSERT_GE(device_count(), kNumReplicas) << "The test requires at least " << kNumReplicas << " devices"; - TF_ASSERT_OK_AND_ASSIGN(auto unoptimized_module, - ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(auto unoptimized_module, + ParseAndReturnVerifiedModule(R"( HloModule m, replica_count=2 e { a = u32[2,2,3] constant({{{0,1,2},{3,4,5}},{{6,7,8},{9,10,11}}}) g = u32[2,4,3] all-gather(a), dimensions={1} })", - kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN(auto executable, - CreateExecutable(std::move(unoptimized_module), - /*run_hlo_passes=*/true)); - TF_ASSERT_OK_AND_ASSIGN(const HloModule* module, - test_runner().HloModuleFromWrapped(executable.get())); + kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto executable, + CreateExecutable(std::move(unoptimized_module), + /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(const HloModule* module, + test_runner().HloModuleFromWrapped(executable.get())); const HloInstruction* root = module->entry_computation()->root_instruction(); EXPECT_THAT(root, @@ -3690,8 +3686,8 @@ TEST_F(CollectiveOpsTestE2E, AllGatherOnChangedDimensionIsCorrect) { EXPECT_THAT(root->fused_expression_root(), GmockMatch(m::Transpose(m::Bitcast(m::Parameter())))); - TF_ASSERT_OK_AND_ASSIGN(std::vector results, - ExecuteReplicated(executable.get(), {{}, {}})); + ASSERT_OK_AND_ASSIGN(std::vector results, + ExecuteReplicated(executable.get(), {{}, {}})); ASSERT_EQ(results.size(), kNumReplicas); Literal expected = LiteralUtil::CreateR3( {{{0, 1, 2}, {3, 4, 5}, {0, 1, 2}, {3, 4, 5}}, @@ -3746,11 +3742,11 @@ TEST_F(CollectiveOpsTestE2E, MultipleModuleDifferentDeviceGroupsShouldRun) { HloModuleConfig config_2 = GetModuleConfigForTest(/*replica_count=*/kNumReplicas_2); - TF_ASSERT_OK_AND_ASSIGN(auto module_1, - ParseAndReturnVerifiedModule(kModuleStr_1, config_1)); + ASSERT_OK_AND_ASSIGN(auto module_1, + ParseAndReturnVerifiedModule(kModuleStr_1, config_1)); - TF_ASSERT_OK_AND_ASSIGN(auto module_2, - ParseAndReturnVerifiedModule(kModuleStr_2, config_2)); + ASSERT_OK_AND_ASSIGN(auto module_2, + ParseAndReturnVerifiedModule(kModuleStr_2, config_2)); int64_t num_elements_1 = ShapeUtil::ElementsIn( module_1->entry_computation()->parameter_instructions()[0]->shape()); @@ -3765,7 +3761,7 @@ TEST_F(CollectiveOpsTestE2E, MultipleModuleDifferentDeviceGroupsShouldRun) { Literal input_literal1_1 = LiteralUtil::CreateFromArray(input1_1); Literal input_literal1_2 = LiteralUtil::CreateFromArray(input1_2); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result_1, ExecuteReplicated(std::move(module_1), std::vector>{ @@ -3783,7 +3779,7 @@ TEST_F(CollectiveOpsTestE2E, MultipleModuleDifferentDeviceGroupsShouldRun) { Literal input_literal2_3 = LiteralUtil::CreateFromArray(input2_3); Literal input_literal2_4 = LiteralUtil::CreateFromArray(input2_4); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result_2, ExecuteReplicated(std::move(module_2), std::vector>{ {&input_literal2_1}, @@ -3823,8 +3819,8 @@ TEST_F(CollectiveOpsTestE2E, CustomCollectiveCallShouldRun) { HloModuleConfig config_1 = GetModuleConfigForTest(/*replica_count=*/kNumReplicas_1); - TF_ASSERT_OK_AND_ASSIGN(auto module_1, - ParseAndReturnVerifiedModule(kModuleStr_1, config_1)); + ASSERT_OK_AND_ASSIGN(auto module_1, + ParseAndReturnVerifiedModule(kModuleStr_1, config_1)); int64_t num_elements_1 = ShapeUtil::ElementsIn( module_1->entry_computation()->parameter_instructions()[0]->shape()); @@ -3836,7 +3832,7 @@ TEST_F(CollectiveOpsTestE2E, CustomCollectiveCallShouldRun) { Literal input_literal1_1 = LiteralUtil::CreateFromArray(input1_1); Literal input_literal1_2 = LiteralUtil::CreateFromArray(input1_2); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result_1, ExecuteReplicated(std::move(module_1), std::vector>{ @@ -3893,14 +3889,14 @@ ENTRY main { HloModuleConfig config = GetModuleConfigForTest(kNumReplicas, kNumPartitions); config.set_debug_options(GetDebugOptionsForTest()); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(hlo_string, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, config)); auto input = LiteralUtil::CreateR1(std::vector(128, 1.0f)); std::vector args = {&input}; std::vector> replica_args(kNumReplicas, args); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult result, - ExecuteReplicated(std::move(module), replica_args)); + ASSERT_OK_AND_ASSIGN(ExecutionResult result, + ExecuteReplicated(std::move(module), replica_args)); for (const auto& literal : result.results) { EXPECT_TRUE(LiteralTestUtil::Near( diff --git a/third_party/xla/xla/backends/gpu/tests/collective_ops_ffi_test.cc b/third_party/xla/xla/backends/gpu/tests/collective_ops_ffi_test.cc index 8cce5f8a4b1236..26d30897d3dbf9 100644 --- a/third_party/xla/xla/backends/gpu/tests/collective_ops_ffi_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/collective_ops_ffi_test.cc @@ -21,6 +21,7 @@ limitations under the License. #include #include +#include #include "absl/base/no_destructor.h" #include "absl/status/status.h" #include "absl/status/status_macros.h" @@ -62,7 +63,6 @@ limitations under the License. #include "xla/stream_executor/stream.h" #include "xla/tests/literal_test_util.h" #include "xla/tsl/platform/errors.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "xla/xla_data.pb.h" @@ -1147,17 +1147,16 @@ TEST_F(CollectiveOpsTestFFI, AllReduce) { } )"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); module->mutable_config() .mutable_debug_options() .set_xla_gpu_executable_num_communication_streams(2); - TF_ASSERT_OK_AND_ASSIGN( - ExecutionResult execution_result, - ExecuteReplicated(std::move(module), - /*arguments=*/std::vector(), - /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), + /*arguments=*/std::vector(), + /*run_hlo_passes=*/false)); absl::Span results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); @@ -1236,14 +1235,13 @@ TEST_P(AllReduceTest, DeviceAllReduce) { )", GetParam()); - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( - ExecutionResult execution_result, - ExecuteReplicated(std::move(module), - /*arguments=*/std::vector(), - /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), + /*arguments=*/std::vector(), + /*run_hlo_passes=*/false)); SynchronizationSignals* signals = global_signals->get(); signals->finished_kernels_counter.Wait(); @@ -1280,14 +1278,13 @@ TEST_P(AllReduceTest, PeerAllReduce) { )", GetParam()); - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( - ExecutionResult execution_result, - ExecuteReplicated(std::move(module), - /*arguments=*/std::vector(), - /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), + /*arguments=*/std::vector(), + /*run_hlo_passes=*/false)); SynchronizationSignals* signals = global_signals->get(); signals->finished_kernels_counter.Wait(); @@ -1325,14 +1322,13 @@ TEST_P(AllReduceTest, MulticastAllReduce) { )", GetParam()); - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( - ExecutionResult execution_result, - ExecuteReplicated(std::move(module), - /*arguments=*/std::vector(), - /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), + /*arguments=*/std::vector(), + /*run_hlo_passes=*/false)); SynchronizationSignals* signals = global_signals->get(); signals->finished_kernels_counter.Wait(); @@ -1369,14 +1365,13 @@ TEST_P(AllReduceTest, SymMulticastAllReduce) { )", GetParam()); - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( - ExecutionResult execution_result, - ExecuteReplicated(std::move(module), - /*arguments=*/std::vector(), - /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), + /*arguments=*/std::vector(), + /*run_hlo_passes=*/false)); SynchronizationSignals* signals = global_signals->get(); signals->finished_kernels_counter.Wait(); @@ -1414,14 +1409,13 @@ TEST_P(AllReduceTest, SymPeerAllReduce) { )", GetParam()); - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( - ExecutionResult execution_result, - ExecuteReplicated(std::move(module), - /*arguments=*/std::vector(), - /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), + /*arguments=*/std::vector(), + /*run_hlo_passes=*/false)); SynchronizationSignals* signals = global_signals->get(); signals->finished_kernels_counter.Wait(); @@ -1469,14 +1463,13 @@ TEST_F(CollectiveOpsTestFFI, DeviceAllReduceWithFrontendAttributes) { } )"; - TF_ASSERT_OK_AND_ASSIGN( - auto module, ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(hlo_string, kNumReplicas)); - TF_ASSERT_OK_AND_ASSIGN( - ExecutionResult execution_result, - ExecuteReplicated(std::move(module), - /*arguments=*/std::vector(), - /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module), + /*arguments=*/std::vector(), + /*run_hlo_passes=*/true)); SynchronizationSignals* signals = global_signals->get(); signals->finished_kernels_counter.Wait(); diff --git a/third_party/xla/xla/backends/gpu/tests/collective_ops_sharded_unsharded_e2e_test.cc b/third_party/xla/xla/backends/gpu/tests/collective_ops_sharded_unsharded_e2e_test.cc index 7a6a0d098efe39..8e0ce4c498af34 100644 --- a/third_party/xla/xla/backends/gpu/tests/collective_ops_sharded_unsharded_e2e_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/collective_ops_sharded_unsharded_e2e_test.cc @@ -20,6 +20,7 @@ limitations under the License. #include #include +#include #include "absl/algorithm/container.h" #include "absl/log/check.h" #include "absl/log/log.h" @@ -37,7 +38,6 @@ limitations under the License. #include "xla/service/hlo_module_config.h" #include "xla/tests/literal_test_util.h" #include "xla/tests/test_utils.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "xla/xla_data.pb.h" #include "tsl/platform/regexp.h" @@ -61,12 +61,12 @@ class CollectiveOpsTestE2EShardedUnsharded : public CollectiveOpsE2ETestBase { << " devices (" << device_count() << " available)"; } - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult ref_execution_result, - ExecuteUnsharded(hlo_text)); + ASSERT_OK_AND_ASSIGN(ExecutionResult ref_execution_result, + ExecuteUnsharded(hlo_text)); const std::vector& ref_results = ref_execution_result.results; ASSERT_EQ(ref_results.size(), 1); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( ExecutionResult execution_result, ExecuteSharded(hlo_text, num_partitions, enable_enzyme_comms_opt)); const std::vector& results = execution_result.results; @@ -190,8 +190,8 @@ class CollectiveOpsTestE2EShardedUnsharded : public CollectiveOpsE2ETestBase { /*replica_count=*/1, /*num_partitions=*/num_partitions); config.mutable_debug_options().set_xla_gpu_enable_triton_gemm(false); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(hlo_text, config)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(hlo_text, config)); auto dimensions = module->entry_computation()->root_instruction()->shape().dimensions(); std::vector root_dims(dimensions.begin(), dimensions.end()); diff --git a/third_party/xla/xla/backends/gpu/tests/collective_pipeline_parallelism_test.cc b/third_party/xla/xla/backends/gpu/tests/collective_pipeline_parallelism_test.cc index 99a130ba461119..ee59c1e17bc7f6 100644 --- a/third_party/xla/xla/backends/gpu/tests/collective_pipeline_parallelism_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/collective_pipeline_parallelism_test.cc @@ -19,6 +19,7 @@ limitations under the License. #include #include +#include #include #include "absl/log/log.h" #include "absl/strings/string_view.h" @@ -32,7 +33,6 @@ limitations under the License. #include "xla/tests/literal_test_util.h" #include "xla/tests/restricted/hlo_test_base_legacy.h" #include "xla/tests/test_utils.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla.pb.h" namespace xla { @@ -121,8 +121,8 @@ TEST_P(CollectivePipelineParallelismTest, HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); std::unique_ptr module; - TF_ASSERT_OK_AND_ASSIGN(module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(module, + ParseAndReturnVerifiedModule(kModuleStr, config)); // Inputs for replica i are // A = {{i+1, i+1}, @@ -139,7 +139,7 @@ TEST_P(CollectivePipelineParallelismTest, for (int64_t i = 0; i < kNumReplicas; ++i) { inputs.push_back({&inputs_a[i], &input_b_replicated}); } - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( std::vector results, ExecuteReplicated(std::move(module), inputs, kNumReplicas, /*run_hlo_passes=*/true)); @@ -316,7 +316,7 @@ TEST_P(CollectivePipelineParallelismTest, NaiveBFSMicrobatch4Replica4) { // Parse HLO module. HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(GetModuleStrWithCommonComputations( /*name=*/"test", kMoreComputationsStr), @@ -345,10 +345,9 @@ TEST_P(CollectivePipelineParallelismTest, NaiveBFSMicrobatch4Replica4) { {&weights_r1, &fake_input}, {&weights_r2, &fake_input}, {&weights_r3, &fake_input}}; - TF_ASSERT_OK_AND_ASSIGN( - std::vector results, - ExecuteReplicated(std::move(module), args, kNumReplicas, - /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(std::vector results, + ExecuteReplicated(std::move(module), args, kNumReplicas, + /*run_hlo_passes=*/true)); // Check pipeline output for last replica. // The combined effect of the pipeline is to scale the input data by 24.0. @@ -440,7 +439,7 @@ TEST_P(CollectivePipelineParallelismTest, NaiveBFSMicrobatch5Replica4) { // Parse HLO module. HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(GetModuleStrWithCommonComputations( /*name=*/"test", kMoreComputationsStr), @@ -473,10 +472,9 @@ TEST_P(CollectivePipelineParallelismTest, NaiveBFSMicrobatch5Replica4) { {&weights_r1, &fake_input}, {&weights_r2, &fake_input}, {&weights_r3, &fake_input}}; - TF_ASSERT_OK_AND_ASSIGN( - std::vector results, - ExecuteReplicated(std::move(module), args, kNumReplicas, - /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(std::vector results, + ExecuteReplicated(std::move(module), args, kNumReplicas, + /*run_hlo_passes=*/true)); EXPECT_TRUE(LiteralTestUtil::NearOrEqual(expected_output, results[3], ErrorSpec{1e-5, 1e-5})); } @@ -563,7 +561,7 @@ TEST_P(CollectivePipelineParallelismTest, // Parse HLO module. HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(GetModuleStrWithCommonComputations( /*name=*/"test", kMoreComputationsStr), @@ -598,10 +596,9 @@ TEST_P(CollectivePipelineParallelismTest, {&weights_r1, &fake_input}, {&weights_r2, &fake_input}, {&weights_r3, &fake_input}}; - TF_ASSERT_OK_AND_ASSIGN( - std::vector results, - ExecuteReplicated(std::move(module), args, kNumReplicas, - /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(std::vector results, + ExecuteReplicated(std::move(module), args, kNumReplicas, + /*run_hlo_passes=*/true)); EXPECT_TRUE(LiteralTestUtil::NearOrEqual(expected_output, results[3], ErrorSpec{1e-5, 1e-5})); } @@ -704,7 +701,7 @@ TEST_P(CollectivePipelineParallelismTest, // Parse HLO module. HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(GetModuleStrWithCommonComputations( /*name=*/"test", kMoreComputationsStr), @@ -739,10 +736,9 @@ TEST_P(CollectivePipelineParallelismTest, {&weights_r1, &fake_input}, {&weights_r2, &fake_input}, {&weights_r3, &fake_input}}; - TF_ASSERT_OK_AND_ASSIGN( - std::vector results, - ExecuteReplicated(std::move(module), args, kNumReplicas, - /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(std::vector results, + ExecuteReplicated(std::move(module), args, kNumReplicas, + /*run_hlo_passes=*/true)); EXPECT_TRUE(LiteralTestUtil::NearOrEqual(expected_output, results[3], ErrorSpec{1e-5, 1e-5})); } @@ -847,7 +843,7 @@ TEST_P(CollectivePipelineParallelismTest, // Parse HLO module. HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(GetModuleStrWithCommonComputations( /*name=*/"test", kMoreComputationsStr), @@ -882,10 +878,9 @@ TEST_P(CollectivePipelineParallelismTest, {&weights_r1, &fake_input}, {&weights_r2, &fake_input}, {&weights_r3, &fake_input}}; - TF_ASSERT_OK_AND_ASSIGN( - std::vector results, - ExecuteReplicated(std::move(module), args, kNumReplicas, - /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(std::vector results, + ExecuteReplicated(std::move(module), args, kNumReplicas, + /*run_hlo_passes=*/true)); EXPECT_TRUE(LiteralTestUtil::NearOrEqual(expected_output, results[3], ErrorSpec{1e-5, 1e-5})); } @@ -947,8 +942,8 @@ TEST_P(CollectivePipelineParallelismTest, SendRecvLoop) { HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); std::unique_ptr module; - TF_ASSERT_OK_AND_ASSIGN(module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(module, + ParseAndReturnVerifiedModule(kModuleStr, config)); // Create input data. std::vector literals; @@ -969,7 +964,7 @@ TEST_P(CollectivePipelineParallelismTest, SendRecvLoop) { } // Execute and check results. - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( std::vector results, ExecuteReplicated(std::move(module), inputs, /*num_replicas=*/kNumPartitions, @@ -1038,8 +1033,8 @@ TEST_P(CollectivePipelineParallelismTest, SendRecvLoop2Devices) { HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); std::unique_ptr module; - TF_ASSERT_OK_AND_ASSIGN(module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(module, + ParseAndReturnVerifiedModule(kModuleStr, config)); // Create input data. std::vector literals; @@ -1060,7 +1055,7 @@ TEST_P(CollectivePipelineParallelismTest, SendRecvLoop2Devices) { } // Execute and check results. - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( std::vector results, ExecuteReplicated(std::move(module), inputs, /*num_replicas=*/kNumPartitions, @@ -1139,8 +1134,8 @@ TEST_P(CollectivePipelineParallelismTest, PartiallyPipelinedAsyncSendRecvLoop) { HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); std::unique_ptr module; - TF_ASSERT_OK_AND_ASSIGN(module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(module, + ParseAndReturnVerifiedModule(kModuleStr, config)); // Create input data. std::vector literals; @@ -1161,7 +1156,7 @@ TEST_P(CollectivePipelineParallelismTest, PartiallyPipelinedAsyncSendRecvLoop) { } // Execute and check results. - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( std::vector results, ExecuteReplicated(std::move(module), inputs, /*num_replicas=*/kNumPartitions, @@ -1242,8 +1237,8 @@ TEST_P(CollectivePipelineParallelismTest, HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); std::unique_ptr module; - TF_ASSERT_OK_AND_ASSIGN(module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(module, + ParseAndReturnVerifiedModule(kModuleStr, config)); // Create input data. std::vector literals; @@ -1264,7 +1259,7 @@ TEST_P(CollectivePipelineParallelismTest, } // Execute and check results. - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( std::vector results, ExecuteReplicated(std::move(module), inputs, /*num_replicas=*/kNumPartitions, @@ -1472,7 +1467,7 @@ TEST_P(CollectivePipelineParallelismTest, HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( std::unique_ptr module, ParseAndReturnVerifiedModule(GetModuleStrWithCommonComputations( /*name=*/"test", kMoreComputationsStr), @@ -1498,10 +1493,9 @@ TEST_P(CollectivePipelineParallelismTest, {&weights_r2, &fake_input}, {&weights_r3, &fake_input}}; // TODO(rosiezou): enable send/recv combiner pass. - TF_ASSERT_OK_AND_ASSIGN( - std::vector results, - ExecuteReplicated(std::move(module), args, kNumReplicas, - /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(std::vector results, + ExecuteReplicated(std::move(module), args, kNumReplicas, + /*run_hlo_passes=*/true)); EXPECT_TRUE(LiteralTestUtil::NearOrEqual(expected_output, results[3], ErrorSpec{1e-5, 1e-5})); } @@ -1659,7 +1653,7 @@ TEST_P(CollectivePipelineParallelismTest, HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto module, ParseAndReturnVerifiedModule(GetModuleStrWithCommonComputations( /*name=*/"test", kMoreComputationsStr), @@ -1684,10 +1678,9 @@ TEST_P(CollectivePipelineParallelismTest, {&weights_r1, &fake_input}, {&weights_r2, &fake_input}, {&weights_r3, &fake_input}}; - TF_ASSERT_OK_AND_ASSIGN( - std::vector results, - ExecuteReplicated(std::move(module), args, kNumReplicas, - /*run_hlo_passes=*/true)); + ASSERT_OK_AND_ASSIGN(std::vector results, + ExecuteReplicated(std::move(module), args, kNumReplicas, + /*run_hlo_passes=*/true)); EXPECT_TRUE(LiteralTestUtil::NearOrEqual( expected_output, results[3], ErrorSpec{/*abs_error=*/1e-5, /*rel_error=*/1e-5})); @@ -2001,8 +1994,8 @@ ENTRY %main.204 (Arg_0.1: f32[4,4096,4096], Arg_1.2: f32[4,5,4096,8192]) HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas, /*num_partitions=*/kNumPartitions); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); // Create device assignment running across partitions. DeviceAssignment device_assignment(/*replica_count=*/kNumReplicas, @@ -2011,13 +2004,13 @@ ENTRY %main.204 (Arg_0.1: f32[4,4096,4096], Arg_1.2: f32[4,5,4096,8192]) device_assignment(0, i) = i; } - TF_ASSERT_OK_AND_ASSIGN(std::vector fake_args, - MakeFakeArguments(module.get())); + ASSERT_OK_AND_ASSIGN(std::vector fake_args, + MakeFakeArguments(module.get())); std::vector args; for (auto& arg : fake_args) { args.push_back(&arg); } - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( std::vector results, ExecuteReplicated(std::move(module), args, /*num_replicas=*/kNumPartitions, &device_assignment, diff --git a/third_party/xla/xla/backends/gpu/tests/gpu_atomic_test.cc b/third_party/xla/xla/backends/gpu/tests/gpu_atomic_test.cc index b1f02fd1787239..f22ba6a6579038 100644 --- a/third_party/xla/xla/backends/gpu/tests/gpu_atomic_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/gpu_atomic_test.cc @@ -13,10 +13,10 @@ See the License for the specific language governing permissions and limitations under the License. ==============================================================================*/ +#include #include #include "xla/backends/gpu/tests/gpu_pjrt_codegen_test.h" #include "xla/stream_executor/cuda/cuda_compute_capability.h" -#include "xla/tsl/lib/core/status_test_util.h" namespace xla { namespace gpu { @@ -46,7 +46,7 @@ TEST_F(GpuAtomicTest, TestStore) { } )"; - TF_ASSERT_OK(CompileAndVerifyIr(hlo_string, R"( + ASSERT_OK(CompileAndVerifyIr(hlo_string, R"( CHECK: store atomic{{.*}}unordered, align 4 )")); } @@ -73,7 +73,7 @@ TEST_F(GpuAtomicTest, TestStoreNoAtomic) { } )"; - TF_ASSERT_OK(CompileAndVerifyIr(hlo_string, R"( + ASSERT_OK(CompileAndVerifyIr(hlo_string, R"( CHECK-NOT: store atomic{{.*}}unordered, align 4 )")); } @@ -101,10 +101,10 @@ TEST_F(GpuAtomicTest, TestAddAtomicF32) { } )"; - TF_ASSERT_OK(CompileAndVerifyIr(hlo_string, IsBuiltWithRocm() ? R"( + ASSERT_OK(CompileAndVerifyIr(hlo_string, IsBuiltWithRocm() ? R"( CHECK: atomicrmw fadd ptr %[[ADDR:.*]], float %[[VALUE:.*]] syncscope("agent-one-as") monotonic )" - : R"( + : R"( CHECK: atomicrmw fadd ptr %[[ADDR:.*]], float %[[VALUE:.*]] monotonic )")); } @@ -138,7 +138,7 @@ TEST_F(GpuAtomicTest, TestAddAtomicF64) { } )"; - TF_ASSERT_OK(CompileAndVerifyIr(hlo_string, R"( + ASSERT_OK(CompileAndVerifyIr(hlo_string, R"( CHECK: atomicrmw fadd ptr %[[ADDR:.*]], double %[[VALUE:.*]] monotonic )")); } diff --git a/third_party/xla/xla/backends/gpu/tests/gpu_spmd_e2e_compile_test.cc b/third_party/xla/xla/backends/gpu/tests/gpu_spmd_e2e_compile_test.cc index f09f5878c38568..1f0043f42f68ee 100644 --- a/third_party/xla/xla/backends/gpu/tests/gpu_spmd_e2e_compile_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/gpu_spmd_e2e_compile_test.cc @@ -28,7 +28,6 @@ limitations under the License. #include "xla/service/hlo_module_config.h" #include "xla/service/hlo_runner_interface.h" #include "xla/tests/hlo_pjrt_test_base.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/xla.pb.h" namespace xla::gpu { @@ -69,7 +68,7 @@ ENTRY entry { absl::StatusOr> executable = CreateExecutable(std::move(hlo_module), /*run_hlo_passes=*/true); - TF_EXPECT_OK(executable.status()); + EXPECT_OK(executable.status()); } TEST_F(GpuSpmdE2ECompileTest, DotSharding) { diff --git a/third_party/xla/xla/backends/gpu/tests/gpu_triton_custom_call_test.cc b/third_party/xla/xla/backends/gpu/tests/gpu_triton_custom_call_test.cc index 33b3e359b8fcf3..37865579f4b4dc 100644 --- a/third_party/xla/xla/backends/gpu/tests/gpu_triton_custom_call_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/gpu_triton_custom_call_test.cc @@ -42,7 +42,6 @@ limitations under the License. #include "xla/stream_executor/cuda/cuda_compute_capability.h" #include "xla/stream_executor/device_description.h" #include "xla/tests/literal_test_util.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/xla_data.pb.h" namespace xla { @@ -329,7 +328,7 @@ TEST_F(GpuIrEmitterUnnestedTest, RunTritonCustomCallWithDeviceSideTMA) { // Run on GPU. absl::StatusOr result_status = Execute(std::move(module), {&input_literal}); - TF_ASSERT_OK(result_status.status()); + ASSERT_OK(result_status.status()); std::vector results = result_status->DecomposeTuple(); EXPECT_TRUE(LiteralTestUtil::Equal(input_literal, results.at(0))); diff --git a/third_party/xla/xla/backends/gpu/tests/multioutput_fusion_test.cc b/third_party/xla/xla/backends/gpu/tests/multioutput_fusion_test.cc index dbaf1d73c1d299..2f41b73548d22e 100644 --- a/third_party/xla/xla/backends/gpu/tests/multioutput_fusion_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/multioutput_fusion_test.cc @@ -20,6 +20,7 @@ limitations under the License. #include #include "xla/tests/xla_test_backend_predicates.h" +#include #include "absl/log/check.h" #include "absl/strings/str_cat.h" #include "absl/strings/substitute.h" @@ -38,7 +39,6 @@ limitations under the License. #include "xla/tests/hlo_pjrt_test_base.h" #include "xla/tests/literal_test_util.h" #include "xla/tests/pjrt_client_registry.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "xla/xla.pb.h" #include "xla/xla_data.pb.h" @@ -114,8 +114,8 @@ class MultiOutputFusionTest : public HloInterpreterReferenceMixin { Literal expect(ShapeUtil::MakeShapeWithDescendingLayout(F32, {size, size})); expect.PopulateWithValue(size * 1.5f * 3.5f); Literal literal_r0 = LiteralUtil::CreateR0(-9.0f); - TF_ASSERT_OK_AND_ASSIGN( - Literal actual, Execute(std::move(hlo_module), {&literal_r0, &arg1})); + ASSERT_OK_AND_ASSIGN(Literal actual, + Execute(std::move(hlo_module), {&literal_r0, &arg1})); EXPECT_TRUE(LiteralTestUtil::Near(expect, actual, kErrorSpec)); } @@ -178,8 +178,8 @@ class MultiOutputFusionTest : public HloInterpreterReferenceMixin { input1.PopulateWithValue(1.); Literal expect = LiteralUtil::CreateR1({size * 1.5f * 3.5f}); - TF_ASSERT_OK_AND_ASSIGN(Literal actual, - Execute(std::move(hlo_module), {&input0, &input1})); + ASSERT_OK_AND_ASSIGN(Literal actual, + Execute(std::move(hlo_module), {&input0, &input1})); EXPECT_TRUE(LiteralTestUtil::Near(expect, actual, kErrorSpec)); } }; @@ -211,8 +211,8 @@ TEST_F(MultiOutputFusionTest, MultiOutputLoopFusion) { })"; auto module = ParseAndReturnVerifiedModule(testcase).value(); auto param = LiteralUtil::CreateR1({1.0, 2.0, 3.0, -1.0}); - TF_ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {¶m}, - /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {¶m}, + /*run_hlo_passes=*/false)); LiteralTestUtil::ExpectR1Equal({0.0, 4.0, 9.0, 1.0}, result); } @@ -239,8 +239,8 @@ TEST_F(MultiOutputFusionTest, MultiOutputLoopFusionBitcastCompatibleShapes) { })"; auto module = ParseAndReturnVerifiedModule(testcase).value(); auto param = LiteralUtil::CreateR1({1.0, 2.0, 3.0, -1.0}); - TF_ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {¶m}, - /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {¶m}, + /*run_hlo_passes=*/false)); LiteralTestUtil::ExpectR1Equal({0.0, 4.0, 9.0, 1.0}, result); } @@ -273,8 +273,8 @@ TEST_F(MultiOutputFusionTest, MultiOutputLoopFeedingMap) { })"; auto module = ParseAndReturnVerifiedModule(testcase).value(); auto param = LiteralUtil::CreateR1({1.0, 2.0, 3.0}); - TF_ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {¶m}, - /*run_hlo_passes=*/false)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {¶m}, + /*run_hlo_passes=*/false)); LiteralTestUtil::ExpectR1Equal({0.0, 4.0, 9.0}, result); } diff --git a/third_party/xla/xla/backends/gpu/tests/nccl_group_execution_test.cc b/third_party/xla/xla/backends/gpu/tests/nccl_group_execution_test.cc index c12ade6c5bc195..0210ed7d399262 100644 --- a/third_party/xla/xla/backends/gpu/tests/nccl_group_execution_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/nccl_group_execution_test.cc @@ -18,6 +18,7 @@ limitations under the License. #include #include +#include #include "absl/strings/string_view.h" #include "absl/types/span.h" #include "xla/hlo/testlib/verified_hlo_module.h" @@ -25,7 +26,6 @@ limitations under the License. #include "xla/service/hlo_module_config.h" #include "xla/tests/hlo_pjrt_test_base.h" #include "xla/tsl/platform/logging.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" namespace xla { @@ -100,9 +100,9 @@ TEST_F(NcclGroupExecutionTest, NcclGroupSendRecvNoWhileLoop) { HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); std::unique_ptr module; - TF_ASSERT_OK_AND_ASSIGN(module, - ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(module, + ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN( std::vector results, ExecuteReplicated(std::move(module), absl::Span{}, kNumReplicas, @@ -146,9 +146,9 @@ TEST_F(NcclGroupExecutionTest, BidirectionalCommunication) { HloModuleConfig config = GetModuleConfigForTest(/*replica_count=*/kNumReplicas); std::unique_ptr module; - TF_ASSERT_OK_AND_ASSIGN(module, - ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(module, + ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN( std::vector results, ExecuteReplicated(std::move(module), absl::Span{}, kNumReplicas, diff --git a/third_party/xla/xla/backends/gpu/tests/nop_custom_call_test.cc b/third_party/xla/xla/backends/gpu/tests/nop_custom_call_test.cc index 00f3b5748fc45b..d4dca98edb4bad 100644 --- a/third_party/xla/xla/backends/gpu/tests/nop_custom_call_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/nop_custom_call_test.cc @@ -16,11 +16,11 @@ limitations under the License. #include #include +#include #include "xla/literal.h" #include "xla/literal_util.h" #include "xla/tests/hlo_pjrt_test_base.h" #include "xla/tests/literal_test_util.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" namespace xla { @@ -49,7 +49,7 @@ TEST_F(NopCustomCallTest, RunAllocateBufferAndUpdate) { })"; auto module = ParseAndReturnVerifiedModule(hlo_text).value(); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( Literal result, Execute(std::move(module), {}, /*run_hlo_passes=*/false)); Literal expected = LiteralUtil::CreateR1({1}); EXPECT_TRUE(LiteralTestUtil::Equal(expected, result)); diff --git a/third_party/xla/xla/backends/gpu/tests/p2p_ops_e2e_test.cc b/third_party/xla/xla/backends/gpu/tests/p2p_ops_e2e_test.cc index cdd40b41241852..66b6229db353bc 100644 --- a/third_party/xla/xla/backends/gpu/tests/p2p_ops_e2e_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/p2p_ops_e2e_test.cc @@ -18,12 +18,12 @@ limitations under the License. #include #include +#include #include "absl/strings/string_view.h" #include "xla/backends/gpu/tests/collective_ops_e2e_test_base.h" #include "xla/literal.h" #include "xla/service/hlo_module_config.h" #include "xla/tests/literal_test_util.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" namespace xla { @@ -59,11 +59,11 @@ TEST_P(P2POps, CollectivePermute) { )"; HloModuleConfig config = GetModuleConfigForTest( /*replica_count=*/kNumReplicas); - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(kModuleStr, config)); + ASSERT_OK_AND_ASSIGN(auto module, + ParseAndReturnVerifiedModule(kModuleStr, config)); - TF_ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, - ExecuteReplicated(std::move(module))); + ASSERT_OK_AND_ASSIGN(ExecutionResult execution_result, + ExecuteReplicated(std::move(module))); const std::vector& results = execution_result.results; ASSERT_EQ(results.size(), kNumReplicas); diff --git a/third_party/xla/xla/backends/gpu/tests/ptx_kernel_test.cc b/third_party/xla/xla/backends/gpu/tests/ptx_kernel_test.cc index 58be44ccf55fbc..45f68cba8e19bc 100644 --- a/third_party/xla/xla/backends/gpu/tests/ptx_kernel_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/ptx_kernel_test.cc @@ -15,11 +15,11 @@ limitations under the License. #include +#include #include #include "absl/strings/string_view.h" #include "xla/literal.h" #include "xla/tests/hlo_pjrt_test_base.h" -#include "xla/tsl/platform/statusor.h" namespace xla { namespace gpu { @@ -47,9 +47,8 @@ TEST_F(PtxKernelE2ETest, ScalarAdd) { }" })"; - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(module_str)); - TF_ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {})); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(module_str)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {})); EXPECT_EQ(result.Get({}), 7.0f); } @@ -73,9 +72,8 @@ TEST_F(PtxKernelE2ETest, TensorAdd) { }" })"; - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(module_str)); - TF_ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {})); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(module_str)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {})); EXPECT_EQ(result.Get({0}), 6.0f); EXPECT_EQ(result.Get({1}), 8.0f); @@ -102,9 +100,8 @@ TEST_F(PtxKernelE2ETest, TensorAddWithoutOutputIndices) { }" })"; - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(module_str)); - TF_ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {})); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(module_str)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {})); EXPECT_EQ(result.Get({0}), 6.0f); EXPECT_EQ(result.Get({1}), 8.0f); @@ -132,9 +129,8 @@ TEST_F(PtxKernelE2ETest, TensorAddWithNonTrivialOutputIndices) { }" })"; - TF_ASSERT_OK_AND_ASSIGN(auto module, - ParseAndReturnVerifiedModule(module_str)); - TF_ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {})); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(module_str)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), {})); EXPECT_EQ(result.Get({0}), 6.0f); EXPECT_EQ(result.Get({1}), 8.0f); diff --git a/third_party/xla/xla/backends/gpu/tests/ragged_dot_test.cc b/third_party/xla/xla/backends/gpu/tests/ragged_dot_test.cc index 8cf1eaf1375ebe..f90118031792c0 100644 --- a/third_party/xla/xla/backends/gpu/tests/ragged_dot_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/ragged_dot_test.cc @@ -16,13 +16,13 @@ limitations under the License. #include #include +#include #include #include "xla/error_spec.h" #include "xla/literal_util.h" #include "xla/tests/hlo_pjrt_interpreter_reference_mixin.h" #include "xla/tests/hlo_pjrt_test_base.h" #include "xla/tests/test_utils.h" -#include "xla/tsl/platform/statusor.h" namespace xla { namespace gpu { @@ -43,11 +43,11 @@ ENTRY main { lhs_ragged_dims={0}, rhs_group_dims={0} } )"; - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_text)); FakeArgumentsOptions options; options.max_bits_of_precision = 10; - TF_ASSERT_OK_AND_ASSIGN(auto fake_arguments, - MakeFakeArguments(module.get(), options)); + ASSERT_OK_AND_ASSIGN(auto fake_arguments, + MakeFakeArguments(module.get(), options)); // Set group sizes to reasonable numbers for ragged_dim_size=6. fake_arguments[2] = LiteralUtil::CreateR1({1, 2, 3}); EXPECT_TRUE(RunAndCompare(std::move(module), @@ -68,11 +68,11 @@ TEST_F(RaggedDotTest, NonContractingWithBatchDims) { lhs_batch_dims={0}, rhs_batch_dims={0}, lhs_ragged_dims={1}, rhs_group_dims={1} })"; - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_text)); FakeArgumentsOptions options; options.max_bits_of_precision = 10; - TF_ASSERT_OK_AND_ASSIGN(auto fake_arguments, - MakeFakeArguments(module.get(), options)); + ASSERT_OK_AND_ASSIGN(auto fake_arguments, + MakeFakeArguments(module.get(), options)); // Set group sizes to reasonable numbers for ragged_dim_size=9. fake_arguments[2] = LiteralUtil::CreateR2({{4, 5}, {7, 2}, {6, 3}}); EXPECT_TRUE(RunAndCompare(std::move(module), @@ -93,11 +93,11 @@ ENTRY main { lhs_ragged_dims={0}, rhs_group_dims={0} } )"; - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_text)); FakeArgumentsOptions options; options.max_bits_of_precision = 10; - TF_ASSERT_OK_AND_ASSIGN(auto fake_arguments, - MakeFakeArguments(module.get(), options)); + ASSERT_OK_AND_ASSIGN(auto fake_arguments, + MakeFakeArguments(module.get(), options)); // Set group sizes to reasonable numbers for ragged_dim_size=6. fake_arguments[2] = LiteralUtil::CreateR1({4, 2}); EXPECT_TRUE(RunAndCompare(std::move(module), @@ -118,11 +118,11 @@ ENTRY main { lhs_ragged_dims={1}, rhs_group_dims={0} } )"; - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo_text)); FakeArgumentsOptions options; options.max_bits_of_precision = 10; - TF_ASSERT_OK_AND_ASSIGN(auto fake_arguments, - MakeFakeArguments(module.get(), options)); + ASSERT_OK_AND_ASSIGN(auto fake_arguments, + MakeFakeArguments(module.get(), options)); // Set group sizes to reasonable numbers for ragged_dim_size=6. fake_arguments[2] = LiteralUtil::CreateR2({{1, 2, 3}, {3, 2, 1}}); EXPECT_TRUE(RunAndCompare(std::move(module), diff --git a/third_party/xla/xla/backends/gpu/tests/replicated_io_feed_test.cc b/third_party/xla/xla/backends/gpu/tests/replicated_io_feed_test.cc index 060bb267a42476..1d6a533633da1b 100644 --- a/third_party/xla/xla/backends/gpu/tests/replicated_io_feed_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/replicated_io_feed_test.cc @@ -18,6 +18,7 @@ limitations under the License. #include #include +#include #include "absl/strings/string_view.h" #include "xla/hlo/ir/hlo_module.h" #include "xla/hlo/testlib/test.h" @@ -29,9 +30,7 @@ limitations under the License. #include "xla/shape_util.h" #include "xla/tests/hlo_pjrt_test_base.h" #include "xla/tests/literal_test_util.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/tsl/platform/logging.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "xla/xla_data.pb.h" @@ -79,12 +78,12 @@ TEST_F(ReplicatedIOFeedTest, InfeedAndOutfeed) { DeviceAssignment device_assn(/*replica_count=*/kNumReplicas, /*computation_count=*/1); device_assn.FillIota(0); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule( - kHloText, GetModuleConfigForTest(kNumReplicas))); - TF_ASSERT_OK(test_runner() - .ExecuteReplicated(std::move(module), opts, &device_assn) - .status()); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule( + kHloText, GetModuleConfigForTest(kNumReplicas))); + ASSERT_OK(test_runner() + .ExecuteReplicated(std::move(module), opts, &device_assn) + .status()); // Verify that each infeed and outfeed is routed correctly. Each replica // should produce 10*replica (indeed) + replica (from HLO) diff --git a/third_party/xla/xla/backends/gpu/tests/simple_optimization_test.cc b/third_party/xla/xla/backends/gpu/tests/simple_optimization_test.cc index 4e460f64116d7a..8386f8cb403f38 100644 --- a/third_party/xla/xla/backends/gpu/tests/simple_optimization_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/simple_optimization_test.cc @@ -16,11 +16,10 @@ limitations under the License. #include #include +#include #include "absl/strings/string_view.h" #include "xla/hlo/testlib/verified_hlo_module.h" #include "xla/tests/hlo_pjrt_test_base.h" -#include "xla/tsl/lib/core/status_test_util.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" namespace xla { @@ -41,11 +40,11 @@ ENTRY e { lhs_contracting_dims={2,3}, rhs_contracting_dims={1,2} })"; - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(kHloText)); - TF_EXPECT_OK(test_runner() - .CreateExecutable(std::move(module), /*run_hlo_passes=*/true) - .status()); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(kHloText)); + EXPECT_OK(test_runner() + .CreateExecutable(std::move(module), /*run_hlo_passes=*/true) + .status()); } } // namespace diff --git a/third_party/xla/xla/backends/gpu/tests/sorting_test.cc b/third_party/xla/xla/backends/gpu/tests/sorting_test.cc index 07a81bc3b821ea..83d066bcb02f4a 100644 --- a/third_party/xla/xla/backends/gpu/tests/sorting_test.cc +++ b/third_party/xla/xla/backends/gpu/tests/sorting_test.cc @@ -20,8 +20,10 @@ limitations under the License. #include #include +#include #include #include "absl/log/check.h" +#include "absl/status/statusor.h" #include "absl/strings/ascii.h" #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" @@ -41,7 +43,6 @@ limitations under the License. #include "xla/shape_util.h" #include "xla/tests/hlo_pjrt_interpreter_reference_mixin.h" #include "xla/tests/hlo_pjrt_test_base.h" -#include "xla/tsl/platform/statusor.h" #include "xla/types.h" #include "xla/xla_data.pb.h" @@ -232,8 +233,8 @@ ENTRY TestComputation { } )"; - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(hlo_text)); HloInstruction* values = module->entry_computation()->GetInstructionWithName("values"); @@ -310,9 +311,9 @@ ENTRY %main { kRadixSortTestSize * 10, // added scratch buffer size ascending ? "false" : "true"); - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); std::vector literals = {std::get<0>(GetParam()).get()}; - TF_ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), literals)); + ASSERT_OK_AND_ASSIGN(Literal result, Execute(std::move(module), literals)); bool has_diff = false; for (int i = 1; i < kRadixSortTestSize; ++i) { @@ -349,10 +350,10 @@ ENTRY %main { kRadixSortTestSize * 20, // added scratch buffer size ascending ? "false" : "true"); - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(hlo)); std::vector literals = {std::get<0>(GetParam()).get()}; - TF_ASSERT_OK_AND_ASSIGN(Literal result_tuple, - Execute(std::move(module), literals)); + ASSERT_OK_AND_ASSIGN(Literal result_tuple, + Execute(std::move(module), literals)); std::vector result = std::move(result_tuple).DecomposeTuple(); bool has_diff = false; From e1e8f864af3c96bf0ae848ed9293904e90638e04 Mon Sep 17 00:00:00 2001 From: Mohammed Anany Date: Wed, 30 Sep 2026 05:06:26 -0700 Subject: [PATCH 57/58] Migrate deprecated TF assertion macros in XLA:GPU codegen and Triton PiperOrigin-RevId: 990917009 --- .../xla/xla/backends/gpu/codegen/BUILD | 1 - .../cubin_custom_kernel_compiler_test.cc | 14 +- .../xla/backends/gpu/codegen/emitters/BUILD | 1 - .../emitters/mlir_kernel_emitter_test.cc | 31 +- .../xla/backends/gpu/codegen/kernels/BUILD | 2 - .../gpu/codegen/kernels/custom_kernel_test.cc | 13 +- .../codegen/kernels/ptx_custom_kernel_test.cc | 43 +- .../gpu/codegen/tools/gpu_test_correctness.cc | 34 +- .../xla/xla/backends/gpu/codegen/triton/BUILD | 9 +- .../gpu/codegen/triton/dot_algorithms_test.cc | 111 +++-- .../gpu/codegen/triton/fusion_test.cc | 11 +- .../gpu/codegen/triton/lowering_util_test.cc | 9 +- .../gpu/codegen/triton/support_legacy_test.cc | 51 ++- .../gpu/codegen/triton/support_test.cc | 417 +++++++++--------- .../backends/gpu/codegen/triton/tests/BUILD | 5 +- .../tests/fusion_emitter_device_test.cc | 15 +- .../triton/tests/scaled_dot_device_test.cc | 4 +- .../triton/tests/triton_test_correctness.cc | 6 +- .../gpu/codegen/triton/tma_utils_test.cc | 15 +- .../codegen/triton/triton_gemm_fusion_test.cc | 11 +- 20 files changed, 381 insertions(+), 422 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/codegen/BUILD b/third_party/xla/xla/backends/gpu/codegen/BUILD index 52bfcd607cb8e3..54beca4b76a892 100644 --- a/third_party/xla/xla/backends/gpu/codegen/BUILD +++ b/third_party/xla/xla/backends/gpu/codegen/BUILD @@ -349,7 +349,6 @@ xla_cc_test( "//xla/service/gpu:launch_dimensions", "//xla/service/gpu:target_constants", "//xla/stream_executor:device_description", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "//xla/tsl/util/proto:parse_text_proto", "@com_google_absl//absl/status:status_matchers", diff --git a/third_party/xla/xla/backends/gpu/codegen/cubin_custom_kernel_compiler_test.cc b/third_party/xla/xla/backends/gpu/codegen/cubin_custom_kernel_compiler_test.cc index ec90e0935a4124..6bf74919b0bdb4 100644 --- a/third_party/xla/xla/backends/gpu/codegen/cubin_custom_kernel_compiler_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/cubin_custom_kernel_compiler_test.cc @@ -20,6 +20,7 @@ limitations under the License. #include #include +#include #include #include "absl/status/status_matchers.h" #include "absl/status/statusor.h" @@ -46,7 +47,6 @@ limitations under the License. #include "xla/service/gpu/launch_dimensions.h" #include "xla/service/gpu/target_constants.h" #include "xla/stream_executor/device_description.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "xla/tsl/util/proto/parse_text_proto.h" #include "xla/xla.pb.h" @@ -99,12 +99,12 @@ TEST(CubinCustomKernelCompilerTest, CallbackInvoked) { TEST_F(HloHardwareIndependentTestBase, TritonCompile) { ObjectPool> mlir_context_pool( []() { return CreateMlirContext(); }); - TF_ASSERT_OK_AND_ASSIGN(BorrowedMlirContext borrowed_context, - mlir_context_pool.GetOrCreate()); + ASSERT_OK_AND_ASSIGN(BorrowedMlirContext borrowed_context, + mlir_context_pool.GetOrCreate()); LoadMlirDialectsForTriton(**borrowed_context); - TF_ASSERT_OK_AND_ASSIGN(mlir::OwningOpRef module, - ParseMlirModuleString(R"( + ASSERT_OK_AND_ASSIGN(mlir::OwningOpRef module, + ParseMlirModuleString(R"( module { xtile.entry_func @random_name(%arg0: memref<125x127xf32>, %arg1: memref<125x127xf32>, %arg2: index) attributes {num_opaque_args = 0 : i32} { %c0 = arith.constant 0 : index @@ -115,7 +115,7 @@ module { xtile.return } })", - **borrowed_context)); + **borrowed_context)); TritonKernelSource triton_source(std::move(module)); auto llvm_compiler = @@ -139,7 +139,7 @@ module { num_ctas: 1 num_stages: 1 )pb")); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TritonWrapperResult result, kernel_compiler .CompileTritonToLlvm("random_name", HloModule{"test_module", {}}, diff --git a/third_party/xla/xla/backends/gpu/codegen/emitters/BUILD b/third_party/xla/xla/backends/gpu/codegen/emitters/BUILD index 241c0f37aad122..35dfc15407d8d9 100644 --- a/third_party/xla/xla/backends/gpu/codegen/emitters/BUILD +++ b/third_party/xla/xla/backends/gpu/codegen/emitters/BUILD @@ -199,7 +199,6 @@ xla_cc_test( "//xla/service/gpu:launch_dimensions", "//xla/stream_executor:device_description", "//xla/tests:xla_internal_test_main", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/status", "@com_google_absl//absl/strings", "@com_google_absl//absl/strings:string_view", diff --git a/third_party/xla/xla/backends/gpu/codegen/emitters/mlir_kernel_emitter_test.cc b/third_party/xla/xla/backends/gpu/codegen/emitters/mlir_kernel_emitter_test.cc index 1f7689274a106a..2dbbbecc330f19 100644 --- a/third_party/xla/xla/backends/gpu/codegen/emitters/mlir_kernel_emitter_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/emitters/mlir_kernel_emitter_test.cc @@ -21,6 +21,7 @@ limitations under the License. #include #include +#include #include #include "absl/status/status.h" #include "absl/strings/str_replace.h" @@ -52,7 +53,6 @@ limitations under the License. #include "xla/service/gpu/gpu_device_info_for_tests.h" #include "xla/service/gpu/launch_dimensions.h" #include "xla/stream_executor/device_description.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla.pb.h" namespace xla { @@ -117,20 +117,19 @@ constexpr absl::string_view kModule = R"( TEST_F(MlirKernelFusionTest, CreateMlirModule) { auto module = ParseAndReturnVerifiedModule(kModule).value(); DummyCopyEmitter emitter; - TF_ASSERT_OK_AND_ASSIGN( - auto mlir_module, - emitter.CreateMLIRModule( - mlir_context_, - *Cast( - module->entry_computation()->root_instruction()), - "fusion", - /*buffer_assignment=*/nullptr)); + ASSERT_OK_AND_ASSIGN(auto mlir_module, + emitter.CreateMLIRModule( + mlir_context_, + *Cast( + module->entry_computation()->root_instruction()), + "fusion", + /*buffer_assignment=*/nullptr)); std::string out; llvm::raw_string_ostream stream(out); stream << *mlir_module; - TF_ASSERT_OK_AND_ASSIGN(auto filecheck_result, RunFileCheck(out, R"( + ASSERT_OK_AND_ASSIGN(auto filecheck_result, RunFileCheck(out, R"( // CHECK: func.func @fusion( // CHECK-SAME: %[[IN:.*]]: tensor<100xf32> {xla.slice_index = 0 // CHECK-SAME: %[[OUT:.*]]: tensor<100xf32> {xla.slice_index = 1 @@ -144,8 +143,8 @@ TEST_F(MlirKernelFusionTest, CreateMlirModule) { } TEST_F(MlirKernelFusionTest, CreateLLVMModule) { - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - ParseAndReturnVerifiedModule(kModule)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + ParseAndReturnVerifiedModule(kModule)); CubinCustomKernelCompiler kernel_compiler( [](llvm::Module& llvm_module, const se::DeviceDescription& descr, const DebugOptions& opts) { return std::vector{}; }, @@ -154,9 +153,9 @@ TEST_F(MlirKernelFusionTest, CreateLLVMModule) { ObjectPool> mlir_context_pool( []() { return CreateMlirContext(); }); MlirKernelFusion emitter(std::make_unique()); - TF_ASSERT_OK_AND_ASSIGN(BorrowedMlirContext borrowed_context, - mlir_context_pool.GetOrCreate()); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(BorrowedMlirContext borrowed_context, + mlir_context_pool.GetOrCreate()); + ASSERT_OK_AND_ASSIGN( LlvmKernelSource source, emitter .CreateLLVMModule( @@ -173,7 +172,7 @@ TEST_F(MlirKernelFusionTest, CreateLLVMModule) { llvm::raw_string_ostream stream(out); stream << *llvm_module.getModuleUnlocked(); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( auto filecheck_result, RunFileCheck( out, absl::StrReplaceAll( diff --git a/third_party/xla/xla/backends/gpu/codegen/kernels/BUILD b/third_party/xla/xla/backends/gpu/codegen/kernels/BUILD index 586cbc804bbe00..93275530b6b3f6 100644 --- a/third_party/xla/xla/backends/gpu/codegen/kernels/BUILD +++ b/third_party/xla/xla/backends/gpu/codegen/kernels/BUILD @@ -69,7 +69,6 @@ xla_cc_test( ":custom_kernel", ":custom_kernel_proto_cc", "//xla/stream_executor:launch_dim", - "//xla/tsl/platform:statusor", "//xla/tsl/util/proto:parse_text_proto", "//xla/tsl/util/proto:proto_matchers", "@com_google_absl//absl/base", @@ -113,7 +112,6 @@ xla_test( "//xla/stream_executor:stream", "//xla/stream_executor:stream_executor_h", "//xla/stream_executor/cuda:cuda_platform", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "//xla/tsl/platform:test_main", "@com_google_absl//absl/log:check", diff --git a/third_party/xla/xla/backends/gpu/codegen/kernels/custom_kernel_test.cc b/third_party/xla/xla/backends/gpu/codegen/kernels/custom_kernel_test.cc index 91c59f0886748b..a1b0b00896d559 100644 --- a/third_party/xla/xla/backends/gpu/codegen/kernels/custom_kernel_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/kernels/custom_kernel_test.cc @@ -24,7 +24,6 @@ limitations under the License. #include "absl/strings/string_view.h" #include "xla/backends/gpu/codegen/kernels/custom_kernel.pb.h" #include "xla/stream_executor/launch_dim.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/util/proto/parse_text_proto.h" #include "xla/tsl/util/proto/proto_matchers.h" @@ -45,7 +44,7 @@ TEST(CustomKernelTest, ToProto) { /*arity=*/42), stream_executor::BlockDim(1, 2, 3), stream_executor::ThreadDim(4, 5, 6), /*shared_memory_bytes=*/7); - TF_ASSERT_OK_AND_ASSIGN(CustomKernelProto proto, custom_kernel.ToProto()); + ASSERT_OK_AND_ASSIGN(CustomKernelProto proto, custom_kernel.ToProto()); EXPECT_THAT( proto, tsl::proto_testing::EqualsProto(R"pb( @@ -71,7 +70,7 @@ TEST(CustomKernelTest, ToProtoWithClusterDims) { stream_executor::BlockDim(1, 2, 3), stream_executor::ThreadDim(4, 5, 6), stream_executor::ClusterDim(7, 8, 9), /*shared_memory_bytes=*/10); - TF_ASSERT_OK_AND_ASSIGN(CustomKernelProto proto, custom_kernel.ToProto()); + ASSERT_OK_AND_ASSIGN(CustomKernelProto proto, custom_kernel.ToProto()); EXPECT_THAT( proto, tsl::proto_testing::EqualsProto(R"pb( @@ -106,8 +105,8 @@ TEST(CustomKernelTest, FromProto) { thread_dims { coordinates { x: 4 y: 5 z: 6 } } shared_memory_bytes: 7 )pb"); - TF_ASSERT_OK_AND_ASSIGN(CustomKernel custom_kernel, - CustomKernel::FromProto(proto, StaticSymbolResolver)); + ASSERT_OK_AND_ASSIGN(CustomKernel custom_kernel, + CustomKernel::FromProto(proto, StaticSymbolResolver)); EXPECT_EQ(custom_kernel.name(), "kernel_name"); EXPECT_EQ(custom_kernel.kernel_spec().kernel_name(), "kernel_name_in_spec"); EXPECT_EQ(custom_kernel.kernel_spec().arity(), 42); @@ -136,8 +135,8 @@ TEST(CustomKernelTest, FromProtoWithClusterDims) { cluster_dim { coordinates { x: 7 y: 8 z: 9 } } shared_memory_bytes: 10 )pb"); - TF_ASSERT_OK_AND_ASSIGN(CustomKernel custom_kernel, - CustomKernel::FromProto(proto, StaticSymbolResolver)); + ASSERT_OK_AND_ASSIGN(CustomKernel custom_kernel, + CustomKernel::FromProto(proto, StaticSymbolResolver)); EXPECT_EQ(custom_kernel.name(), "kernel_name"); EXPECT_EQ(custom_kernel.kernel_spec().kernel_name(), "kernel_name_in_spec"); EXPECT_EQ(custom_kernel.kernel_spec().arity(), 42); diff --git a/third_party/xla/xla/backends/gpu/codegen/kernels/ptx_custom_kernel_test.cc b/third_party/xla/xla/backends/gpu/codegen/kernels/ptx_custom_kernel_test.cc index 14e331f7dc5a28..b9ecc47c3c4c61 100644 --- a/third_party/xla/xla/backends/gpu/codegen/kernels/ptx_custom_kernel_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/kernels/ptx_custom_kernel_test.cc @@ -32,7 +32,6 @@ limitations under the License. #include "xla/stream_executor/launch_dim.h" #include "xla/stream_executor/stream.h" #include "xla/stream_executor/stream_executor.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" namespace xla::gpu::kernel { @@ -85,17 +84,17 @@ TEST(PtxCustomKernelTest, GetPtxCustomKernel) { int64_t length = 4; int64_t byte_length = sizeof(int32_t) * length; se::gpu::CudaPlatform platform; - TF_ASSERT_OK_AND_ASSIGN(se::StreamExecutor * executor, - platform.ExecutorForDevice(0)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(se::StreamExecutor * executor, + platform.ExecutorForDevice(0)); + ASSERT_OK_AND_ASSIGN( CustomKernel custom_kernel, GetPtxCustomKernel("AddI32", kAddI32KernelPtx, 3, se::BlockDim(4), se::ThreadDim(1), byte_length)); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr kernel, - executor->LoadKernel(custom_kernel.kernel_spec())); + ASSERT_OK_AND_ASSIGN(std::unique_ptr kernel, + executor->LoadKernel(custom_kernel.kernel_spec())); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr stream, - executor->CreateStream()); + ASSERT_OK_AND_ASSIGN(std::unique_ptr stream, + executor->CreateStream()); se::DeviceAddress a = executor->AllocateArray(length, 0); se::DeviceAddress b = executor->AllocateArray(length, 0); se::DeviceAddress c = executor->AllocateArray(length, 0); @@ -125,17 +124,17 @@ TEST(PtxCustomKernelTest, GetPtxCustomKernelWithClusterDim) { int64_t length = 4; int64_t byte_length = sizeof(int32_t) * length; se::gpu::CudaPlatform platform; - TF_ASSERT_OK_AND_ASSIGN(se::StreamExecutor * executor, - platform.ExecutorForDevice(0)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(se::StreamExecutor * executor, + platform.ExecutorForDevice(0)); + ASSERT_OK_AND_ASSIGN( CustomKernel custom_kernel, GetPtxCustomKernel("AddI32", kAddI32KernelPtx, 3, se::BlockDim(4), se::ThreadDim(1), se::ClusterDim(2), byte_length)); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr kernel, - executor->LoadKernel(custom_kernel.kernel_spec())); + ASSERT_OK_AND_ASSIGN(std::unique_ptr kernel, + executor->LoadKernel(custom_kernel.kernel_spec())); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr stream, - executor->CreateStream()); + ASSERT_OK_AND_ASSIGN(std::unique_ptr stream, + executor->CreateStream()); se::DeviceAddress a = executor->AllocateArray(length, 0); se::DeviceAddress b = executor->AllocateArray(length, 0); se::DeviceAddress c = executor->AllocateArray(length, 0); @@ -203,18 +202,18 @@ TEST(PtxCustomKernelTest, GetOwnedPtxCustomKernel) { int64_t length = 4; int64_t byte_length = sizeof(int32_t) * length; se::gpu::CudaPlatform platform; - TF_ASSERT_OK_AND_ASSIGN(se::StreamExecutor * executor, - platform.ExecutorForDevice(0)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(se::StreamExecutor * executor, + platform.ExecutorForDevice(0)); + ASSERT_OK_AND_ASSIGN( CustomKernel custom_kernel, GetOwnedPtxCustomKernel("AddI32", kAddI32KernelPtx, 3, se::BlockDim(4), se::ThreadDim(1), byte_length)); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr kernel, - executor->LoadKernel(custom_kernel.kernel_spec())); + ASSERT_OK_AND_ASSIGN(std::unique_ptr kernel, + executor->LoadKernel(custom_kernel.kernel_spec())); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr stream, - executor->CreateStream()); + ASSERT_OK_AND_ASSIGN(std::unique_ptr stream, + executor->CreateStream()); se::DeviceAddress a = executor->AllocateArray(length, 0); se::DeviceAddress b = executor->AllocateArray(length, 0); se::DeviceAddress c = executor->AllocateArray(length, 0); diff --git a/third_party/xla/xla/backends/gpu/codegen/tools/gpu_test_correctness.cc b/third_party/xla/xla/backends/gpu/codegen/tools/gpu_test_correctness.cc index e24d35c3b0630f..f48a26c8fd2f5b 100644 --- a/third_party/xla/xla/backends/gpu/codegen/tools/gpu_test_correctness.cc +++ b/third_party/xla/xla/backends/gpu/codegen/tools/gpu_test_correctness.cc @@ -18,6 +18,7 @@ limitations under the License. #include #include +#include #include #include "absl/log/check.h" #include "absl/log/log.h" @@ -41,7 +42,6 @@ limitations under the License. #include "xla/shape.h" #include "xla/tests/hlo_pjrt_interpreter_reference_mixin.h" #include "xla/tests/hlo_pjrt_test_base.h" -#include "xla/tsl/lib/core/status_test_util.h" struct Flags { std::string input_file = ""; @@ -82,7 +82,7 @@ absl::Status TestBijection(const IndexingMap& map, } TEST_F(CorrectnessTest, RunAndCompare) { - TF_ASSERT_OK_AND_ASSIGN(auto module, LoadTestModule(flags.input_file)); + ASSERT_OK_AND_ASSIGN(auto module, LoadTestModule(flags.input_file)); EXPECT_TRUE(RunAndCompareNoHloPasses( std::move(module), ErrorSpec{flags.abs_error_bound, flags.rel_error_bound})); @@ -113,21 +113,21 @@ std::pair> ParseHeroAndIds( TEST_F(CorrectnessTest, InputIndexingIsBijection) { auto mlir_context = GetMlirContextForTest(); RegisterSymbolicExprStorage(&mlir_context); - TF_ASSERT_OK_AND_ASSIGN(auto module, LoadTestModule(flags.input_file)); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr emitter_data, - GetEmitter(*module)); + ASSERT_OK_AND_ASSIGN(auto module, LoadTestModule(flags.input_file)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr emitter_data, + GetEmitter(*module)); for (const auto& [hero_name, ids] : flags.bijection_inputs) { - TF_ASSERT_OK_AND_ASSIGN(int64_t hero_index, - GetHeroIndex(hero_name, *emitter_data->analysis)); + ASSERT_OK_AND_ASSIGN(int64_t hero_index, + GetHeroIndex(hero_name, *emitter_data->analysis)); auto indexing = emitter_data->emitter->ComputeThreadIdToInputIndexing( hero_index, &mlir_context); ASSERT_TRUE(indexing.has_value()); for (int64_t id : ids) { - TF_ASSERT_OK(TestBijection(indexing.value()[id], - emitter_data->analysis->fusion_hero(hero_index) - .GetOperand(id) - .shape() - .dimensions())) + ASSERT_OK(TestBijection(indexing.value()[id], + emitter_data->analysis->fusion_hero(hero_index) + .GetOperand(id) + .shape() + .dimensions())) << "Expected operand " << id << " of " << hero_name << " (root index " << hero_index << ") to be read exactly once."; } @@ -137,15 +137,15 @@ TEST_F(CorrectnessTest, InputIndexingIsBijection) { TEST_F(CorrectnessTest, OutputIndexingIsBijection) { auto mlir_context = GetMlirContextForTest(); RegisterSymbolicExprStorage(&mlir_context); - TF_ASSERT_OK_AND_ASSIGN(auto module, LoadTestModule(flags.input_file)); - TF_ASSERT_OK_AND_ASSIGN(auto emitter_data, GetEmitter(*module)); + ASSERT_OK_AND_ASSIGN(auto module, LoadTestModule(flags.input_file)); + ASSERT_OK_AND_ASSIGN(auto emitter_data, GetEmitter(*module)); for (const auto& hero_name : flags.bijection_outputs) { - TF_ASSERT_OK_AND_ASSIGN(int64_t hero_index, - GetHeroIndex(hero_name, *emitter_data->analysis)); + ASSERT_OK_AND_ASSIGN(int64_t hero_index, + GetHeroIndex(hero_name, *emitter_data->analysis)); auto indexing = emitter_data->emitter->ComputeThreadIdToOutputIndexing( hero_index, &mlir_context); ASSERT_TRUE(indexing.has_value()); - TF_ASSERT_OK(TestBijection( + ASSERT_OK(TestBijection( *indexing, GetFirstArrayShape( emitter_data->analysis->fusion_root(hero_index).shape()) .dimensions())) diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/BUILD b/third_party/xla/xla/backends/gpu/codegen/triton/BUILD index 066f33d49a718b..1444ba29728bb6 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/BUILD +++ b/third_party/xla/xla/backends/gpu/codegen/triton/BUILD @@ -94,7 +94,6 @@ xla_cc_test( "//xla/stream_executor:device_description", "//xla/stream_executor:launch_dim", "//xla/tests:xla_internal_test_main", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_matchers", "@com_google_googletest//:gtest", @@ -143,7 +142,6 @@ xla_cc_test( "//xla/service/llvm_ir:llvm_util", "//xla/stream_executor:launch_dim", "//xla/stream_executor/gpu:tma_metadata", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/log:check", "@com_google_googletest//:gtest_main", "@llvm-project//mlir:ArithDialect", @@ -380,7 +378,6 @@ xla_test( "//xla/stream_executor/cuda:cuda_compute_capability", "//xla/tests:hlo_interpreter_reference_mixin", "//xla/tests:xla_internal_test_main", # fixdeps: keep - "//xla/tsl/lib/core:status_test_util", "//xla/tsl/platform:env", "//xla/tsl/platform:errors", "//xla/tsl/platform:test", @@ -445,7 +442,6 @@ xla_test( "//xla/tests:test_utils", "//xla/tests:xla_internal_test_main", # fixdeps: keep "//xla/tsl/platform:errors", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/algorithm:container", "@com_google_absl//absl/container:flat_hash_map", "@com_google_absl//absl/log", @@ -631,7 +627,6 @@ xla_cc_test( "//xla/stream_executor:device_description", "//xla/stream_executor/cuda:cuda_compute_capability", "//xla/stream_executor/rocm:rocm_compute_capability", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/algorithm:container", "@com_google_absl//absl/container:flat_hash_set", "@com_google_absl//absl/log", @@ -678,11 +673,10 @@ xla_test( "//xla/stream_executor/cuda:cuda_compute_capability", "//xla/tests:hlo_pjrt_interpreter_reference_mixin", "//xla/tests:hlo_pjrt_test_base", - "//xla/tsl/lib/core:status_test_util", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/log:check", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_matchers", + "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", "@com_google_googletest//:gtest_main", ], @@ -883,7 +877,6 @@ xla_cc_test( ":tma_utils", "//xla/backends/gpu/codegen/triton/ir:triton_xla", "//xla/stream_executor/gpu:tma_metadata", - "//xla/tsl/platform:statusor", "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_matchers", "@com_google_googletest//:gtest_main", diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/dot_algorithms_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/dot_algorithms_test.cc index 21d02db5101fc1..84f5222bc11c52 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/dot_algorithms_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/dot_algorithms_test.cc @@ -70,7 +70,6 @@ limitations under the License. #include "xla/tests/hlo_pjrt_interpreter_reference_mixin.h" #include "xla/tests/test_utils.h" #include "xla/tsl/platform/errors.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla.pb.h" #include "xla/xla_data.pb.h" @@ -510,7 +509,7 @@ TEST_F(Triton6xBF16GemmTest, Emit6xBF16GemmWhenBothInputsAreF32) { "num_stages":1,"num_warps":1,"num_ctas":1}}} } )"; - TF_ASSERT_OK( + ASSERT_OK( CreateTritonIrFromHloTextAndFileCheckForDot(kHloText, "triton_dot", R"( CHECK: %[[INFINITY:.*]] = arith.constant dense<0x7F800000> : tensor<32x32xf32> CHECK: %[[C0:.*]] = arith.constant dense<0.000000e+00> : tensor<32x32xf32> @@ -564,7 +563,7 @@ TEST_F(Triton6xBF16GemmTest, Triton6xBF16GemmWorksForLongContractingDimension) { "num_stages":1,"num_warps":4, "num_ctas":1}}} } )"; - TF_ASSERT_OK( + ASSERT_OK( CreateTritonIrFromHloTextAndFileCheckForDot(kHloText, "triton_dot", R"( CHECK-COUNT-6: %{{.*}} = tt.dot %{{.*}}, %{{.*}}, %{{.*}} : tensor<64x32xbf16> * tensor<32x32xbf16> -> tensor<64x32xf32> )")); @@ -587,8 +586,8 @@ TEST_F(Triton6xBF16GemmTest, Emit6xBF16GemmEndToEnd) { algorithm=dot_bf16_bf16_f32_x6 } )"; - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr verified_module, - ParseAndReturnVerifiedModule(kHloText)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr verified_module, + ParseAndReturnVerifiedModule(kHloText)); CompileAndOptionallyVerifyPtx(std::move(verified_module), R"( CHECK: mma.sync.aligned.{{.*}}.row.col.f32.bf16.bf16.f32 @@ -626,7 +625,7 @@ TEST_F(Triton3xBF16GemmTest, Emit3xBF16GemmWhenBothInputsAreF32) { "num_stages":1,"num_warps":1,"num_ctas":1}}} } )"; - TF_ASSERT_OK( + ASSERT_OK( CreateTritonIrFromHloTextAndFileCheckForDot(kHloText, "triton_dot", R"( CHECK: %[[INFINITY:.*]] = arith.constant dense<0x7F800000> : tensor<32x32xf32> CHECK: %[[C0:.*]] = arith.constant dense<0.000000e+00> : tensor<32x32xf32> @@ -674,7 +673,7 @@ TEST_F(Triton3xBF16GemmTest, Triton3xBF16GemmWorksForLongContractingDimension) { "num_stages":1,"num_warps":4, "num_ctas":1}}} } )"; - TF_ASSERT_OK( + ASSERT_OK( CreateTritonIrFromHloTextAndFileCheckForDot(kHloText, "triton_dot", R"( CHECK-COUNT-3: %{{.*}} = tt.dot %{{.*}}, %{{.*}}, %{{.*}} : tensor<64x32xbf16> * tensor<32x32xbf16> -> tensor<64x32xf32> )")); @@ -697,8 +696,8 @@ TEST_F(Triton3xBF16GemmTest, Emit3xBF16GemmEndToEnd) { algorithm=dot_bf16_bf16_f32_x3 } )"; - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr verified_module, - ParseAndReturnVerifiedModule(kHloText)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr verified_module, + ParseAndReturnVerifiedModule(kHloText)); CompileAndOptionallyVerifyPtx(std::move(verified_module), R"( CHECK: mma.sync.aligned.{{.*}}.row.col.f32.bf16.bf16.f32 @@ -726,8 +725,8 @@ TEST_F(TritonAlgorithmTest, Algorithm_BF16_BF16_F32_X3) { )"; constexpr absl::string_view kPattern = R"(CHECK: "kind":"__triton_nested_gemm_fusion")"; - TF_ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); - TF_ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); + ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); + ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); EXPECT_TRUE(ok); } @@ -749,8 +748,8 @@ TEST_F(TritonAlgorithmTest, Algorithm_BF16_BF16_F32_X6) { )"; constexpr absl::string_view kPattern = R"(CHECK: "kind":"__triton_nested_gemm_fusion")"; - TF_ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); - TF_ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); + ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); + ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); EXPECT_TRUE(ok); } @@ -774,8 +773,8 @@ TEST_F(TritonAlgorithmTest, Algorithm_TF32_TF32_F32) { CHECK: algorithm=dot_tf32_tf32_f32 CHECK: "kind":"__triton_nested_gemm_fusion" )"; - TF_ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); - TF_ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); + ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); + ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); EXPECT_TRUE(ok); } @@ -797,8 +796,8 @@ TEST_F(TritonAlgorithmTest, Algorithm_TF32_TF32_F32_X3) { )"; constexpr absl::string_view kPattern = R"(CHECK: "kind":"__triton_nested_gemm_fusion")"; - TF_ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); - TF_ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); + ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); + ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); EXPECT_TRUE(ok); } @@ -823,8 +822,8 @@ TEST_F(TritonAlgorithmTest, Algorithm_BF16_BF16_F32) { )"; constexpr absl::string_view kPattern = R"(CHECK: "kind":"__triton_nested_gemm_fusion")"; - TF_ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); - TF_ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); + ASSERT_OK_AND_ASSIGN(auto module, GetOptimizedModule(kHloText)); + ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module->ToString(), kPattern)); EXPECT_TRUE(ok); } @@ -1149,11 +1148,11 @@ TEST_P(NumericTestsForBlas, Infinity) { GTEST_SKIP() << "hipBLASLt FAST_TF32 inf*1 returns NaN on MI350 "; } std::string hlo_text = HloText(); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - GetOptimizedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + GetOptimizedModule(hlo_text)); auto module_text = module->ToString(); - TF_ASSERT_OK_AND_ASSIGN(auto ok, - RunFileCheck(module_text, kCheckTritionNestedGemm)); + ASSERT_OK_AND_ASSIGN(auto ok, + RunFileCheck(module_text, kCheckTritionNestedGemm)); ASSERT_TRUE(ok); auto reference_module = GetReferenceModuleForCublas(); @@ -1168,11 +1167,11 @@ TEST_P(NumericTestsForBlas, Infinity) { TEST_P(NumericTestsForBlas, NaN) { std::string hlo_text = HloText(); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - GetOptimizedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + GetOptimizedModule(hlo_text)); auto module_text = module->ToString(); - TF_ASSERT_OK_AND_ASSIGN(auto ok, - RunFileCheck(module_text, kCheckTritionNestedGemm)); + ASSERT_OK_AND_ASSIGN(auto ok, + RunFileCheck(module_text, kCheckTritionNestedGemm)); ASSERT_TRUE(ok); auto reference_module = GetReferenceModuleForCublas(); @@ -1187,11 +1186,11 @@ TEST_P(NumericTestsForBlas, NaN) { TEST_P(NumericTestsForBlas, InputsWithLargeExponent) { std::string hlo_text = HloText(); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - GetOptimizedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + GetOptimizedModule(hlo_text)); auto module_text = module->ToString(); - TF_ASSERT_OK_AND_ASSIGN(auto ok, - RunFileCheck(module_text, kCheckTritionNestedGemm)); + ASSERT_OK_AND_ASSIGN(auto ok, + RunFileCheck(module_text, kCheckTritionNestedGemm)); ASSERT_TRUE(ok); auto reference_module = GetReferenceModuleForCublas(); @@ -1208,11 +1207,11 @@ TEST_P(NumericTestsForBlas, InputsWithLargeExponent) { TEST_P(NumericTestsForBlas, PrecisionCheck) { std::string hlo_text = HloText(); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - GetOptimizedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + GetOptimizedModule(hlo_text)); auto module_text = module->ToString(); - TF_ASSERT_OK_AND_ASSIGN(auto ok, - RunFileCheck(module_text, kCheckTritionNestedGemm)); + ASSERT_OK_AND_ASSIGN(auto ok, + RunFileCheck(module_text, kCheckTritionNestedGemm)); ASSERT_TRUE(ok); auto reference_module = GetReferenceModuleForCublas(); @@ -1231,10 +1230,10 @@ TEST_P(NumericTestsForTriton, Infinity) { // It is the tricky cases for X3 and X6 algorithms. They should mask the NaN // intermediate results. std::string hlo_text = HloText(); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - GetOptimizedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + GetOptimizedModule(hlo_text)); auto module_text = module->ToString(); - TF_ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module_text, kPattern)); + ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module_text, kPattern)); ASSERT_TRUE(ok); EXPECT_TRUE(RunAndCompareNoHloPasses(std::move(module), infinity_arguments_ptrs(), @@ -1245,11 +1244,11 @@ TEST_P(NumericTestsForTriton, Infinity) { TEST_P(NumericTestsForTriton, NaN) { std::string hlo_text = HloText(); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - GetOptimizedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + GetOptimizedModule(hlo_text)); auto module_text = module->ToString(); - TF_ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module_text, kPattern)); + ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module_text, kPattern)); ASSERT_TRUE(ok); EXPECT_TRUE(RunAndCompareNoHloPasses(std::move(module), nan_arguments_ptrs(), ErrorSpec{/*aabs=*/0, /*arel=*/0})) @@ -1259,10 +1258,10 @@ TEST_P(NumericTestsForTriton, NaN) { TEST_P(NumericTestsForTriton, InputsWithLargeExponent) { std::string hlo_text = HloText(); - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - GetOptimizedModule(hlo_text)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + GetOptimizedModule(hlo_text)); auto module_text = module->ToString(); - TF_ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module_text, kPattern)); + ASSERT_OK_AND_ASSIGN(auto ok, RunFileCheck(module_text, kPattern)); ASSERT_TRUE(ok); EXPECT_TRUE(RunAndCompareNoHloPasses( @@ -1962,10 +1961,9 @@ TEST_P(PrecisionTests, PrecisionCheck) { constexpr int kLhsOuterDim = 1024; constexpr int kRhsOuterDim = 1024; constexpr int kContractingDim = 8; - TF_ASSERT_OK_AND_ASSIGN( - std::unique_ptr test_module, - GetSimpleDotModule(kLhsOuterDim, kRhsOuterDim, kContractingDim, algorithm, - backend)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr test_module, + GetSimpleDotModule(kLhsOuterDim, kRhsOuterDim, + kContractingDim, algorithm, backend)); FakeArgumentsOptions options; options.max_bits_of_precision = 23; ASSERT_OK_AND_ASSIGN(std::vector fake_arguments, @@ -1976,16 +1974,16 @@ TEST_P(PrecisionTests, PrecisionCheck) { std::vector ref_result = RunReferenceDot(GetLiteralPointers(fake_arguments), kLhsOuterDim, kRhsOuterDim, kContractingDim); - TF_ASSERT_OK_AND_ASSIGN(auto executable, test_runner().CreateExecutable( - std::move(test_module), false)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(auto executable, test_runner().CreateExecutable( + std::move(test_module), false)); + ASSERT_OK_AND_ASSIGN( Literal test_result, test_runner().ExecuteWithExecutable(executable.get(), fake_arguments)); std::vector profile_times; profile_times.reserve(100); for (int i = 0; i < 100; ++i) { auto start = absl::Now(); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( Literal iter_result, test_runner().ExecuteWithExecutable(executable.get(), fake_arguments)); auto elapsed = absl::Now() - start; @@ -2031,7 +2029,7 @@ TEST_P(PrecisionTests, CheckPrecisionDegradationAlongKDimension) { csv_writer.appendRow( {"iterations_along_k", "max(abs(rel_errors))", "std_dev(rel_errors)"}); for (int k = kMinKSize; k <= kMaxKSize; k *= 2) { - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( std::unique_ptr test_module, GetSimpleDotModule(kMSize, kNSize, k, algorithm, backend)); FakeArgumentsOptions options; @@ -2045,10 +2043,9 @@ TEST_P(PrecisionTests, CheckPrecisionDegradationAlongKDimension) { GetLiteralPointers(fake_arguments); std::vector ref_result = RunReferenceDot(fake_argument_ptrs, kMSize, kNSize, k); - TF_ASSERT_OK_AND_ASSIGN( - auto executable, - test_runner().CreateExecutable(std::move(test_module), false)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN(auto executable, test_runner().CreateExecutable( + std::move(test_module), false)); + ASSERT_OK_AND_ASSIGN( Literal test_result, test_runner().ExecuteWithExecutable(executable.get(), fake_arguments)); std::vector rel_errors = diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/fusion_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/fusion_test.cc index 6ef30440bc4b58..2a221c3554668e 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/fusion_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/fusion_test.cc @@ -38,7 +38,6 @@ limitations under the License. #include "xla/service/gpu/target_constants.h" #include "xla/stream_executor/device_description.h" #include "xla/stream_executor/launch_dim.h" -#include "xla/tsl/platform/statusor.h" namespace xla { namespace gpu { @@ -50,7 +49,7 @@ class TritonFusionTest : public HloHardwareIndependentTestBase {}; TEST_F(TritonFusionTest, TritonFusionWithBlockLevelFusionConfig_LaunchConfigIsCorrect) { - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(R"( triton_computation { param_0 = f32[125,127] parameter(0) ROOT abs = f32[125,127] abs(param_0) @@ -90,7 +89,7 @@ ENTRY entry_computation { TEST_F(TritonFusionTest, TritonFusionWithoutBlockLevelFusionConfig_LaunchConfigIsNullopt) { - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(R"( triton_computation { param_0 = f32[125,127] parameter(0) ROOT abs = f32[125,127] abs(param_0) @@ -111,8 +110,8 @@ ENTRY entry_computation { ObjectPool> mlir_context_pool( []() { return CreateMlirContext(); }); - TF_ASSERT_OK_AND_ASSIGN(BorrowedMlirContext borrowed_context, - mlir_context_pool.GetOrCreate()); + ASSERT_OK_AND_ASSIGN(BorrowedMlirContext borrowed_context, + mlir_context_pool.GetOrCreate()); std::unique_ptr emitter = GetFusionEmitter(PreBufferAssignmentFusionInfo{analysis}); @@ -136,7 +135,7 @@ ENTRY entry_computation { TEST_F( TritonFusionTest, TritonFusionWithBlockLevelFusionConfig_LaunchConfigOverrideWorksCorrectly) { - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(R"( + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(R"( triton_computation { param_0 = f32[125,127] parameter(0) ROOT abs = f32[125,127] abs(param_0) diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/lowering_util_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/lowering_util_test.cc index ac21e7aea1b2a0..a5fb20a86a976c 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/lowering_util_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/lowering_util_test.cc @@ -36,7 +36,6 @@ limitations under the License. #include "xla/service/llvm_ir/llvm_util.h" #include "xla/stream_executor/gpu/tma_metadata.h" #include "xla/stream_executor/launch_dim.h" -#include "xla/tsl/platform/statusor.h" #include "triton/Dialect/Triton/IR/Dialect.h" #include "triton/Dialect/TritonGPU/IR/Dialect.h" @@ -94,8 +93,8 @@ module { mlir::OwningOpRef module = ParseModule(kMlirModule); mlir::LLVM::LLVMFuncOp func_op = *module->getOps().begin(); - TF_ASSERT_OK_AND_ASSIGN(stream_executor::gpu::TmaMetadata tma_metadata, - xgt::ExtractTmaMetadata(func_op)); + ASSERT_OK_AND_ASSIGN(stream_executor::gpu::TmaMetadata tma_metadata, + xgt::ExtractTmaMetadata(func_op)); EXPECT_EQ(tma_metadata.arg_index_to_tma_info.size(), 2); EXPECT_TRUE(tma_metadata.arg_index_to_tma_info.contains(1)); @@ -130,8 +129,8 @@ module attributes {ttg.global_scratch_memory_alignment = 1 : i32, ttg.global_scr mlir::OwningOpRef module = ParseModule(kMlirModule); mlir::LLVM::LLVMFuncOp func_op = *module->getOps().begin(); - TF_ASSERT_OK_AND_ASSIGN(stream_executor::ThreadDim thread_dims, - xgt::ExtractThreadDims(module.get(), func_op)); + ASSERT_OK_AND_ASSIGN(stream_executor::ThreadDim thread_dims, + xgt::ExtractThreadDims(module.get(), func_op)); EXPECT_EQ(thread_dims, stream_executor::ThreadDim(32, 1, 1)); } diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/support_legacy_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/support_legacy_test.cc index 13156e1a4a6790..febff3d6c3d57b 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/support_legacy_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/support_legacy_test.cc @@ -25,6 +25,7 @@ limitations under the License. #include "absl/log/check.h" #include "absl/status/status.h" #include "absl/status/status_matchers.h" +#include "absl/status/statusor.h" #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" #include "absl/strings/substitute.h" @@ -48,8 +49,6 @@ limitations under the License. #include "xla/stream_executor/device_description.h" #include "xla/tests/hlo_pjrt_interpreter_reference_mixin.h" #include "xla/tests/hlo_pjrt_test_base.h" -#include "xla/tsl/lib/core/status_test_util.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla.pb.h" #include "xla/xla_data.pb.h" @@ -131,12 +130,12 @@ ENTRY e { })"; const std::string hlo_test = absl::Substitute( kHloTestTemplate, lhs, rhs, output, HloOpcodeString(opcode)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(hlo_test, /*data_type=*/{}, opcode)); if (legacy_triton::IsTritonSupportedInstruction(ti.Instruction(), GetComputeCapability())) { - TF_EXPECT_OK( + EXPECT_OK( ApplyFloatNormalization(ti.Module().get(), GetComputeCapability())); EXPECT_TRUE(RunAndCompareNoHloPasses( std::move(ti.Module()), @@ -284,9 +283,9 @@ ENTRY e { param.is_the_majormost_dim_being_sliced ? 1 : 0, // start_index0 param.is_the_majormost_dim_being_sliced ? 0 : 1 // start_index1 ); - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction dot, - ParseTemplateAndGetInstruction( - hlo_test, /*data_type=*/{}, HloOpcode::kDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction dot, + ParseTemplateAndGetInstruction( + hlo_test, /*data_type=*/{}, HloOpcode::kDot)); HloInstruction* dynamic_slice = FindInstruction(dot.Module().get(), HloOpcode::kDynamicSlice); ASSERT_NE(dynamic_slice, nullptr); @@ -304,7 +303,7 @@ ENTRY e { if (is_supported_instruction) { // TODO(goncharov): Change to `EXPECT_FALSE(is_supported_instruction)`. GTEST_SKIP() << "The generic emitter does not support dynamic slice yet."; - TF_EXPECT_OK( + EXPECT_OK( ApplyFloatNormalization(dot.Module().get(), GetComputeCapability())); EXPECT_TRUE(RunAndCompareNoHloPasses( std::move(dot.Module()), ErrorSpec{/*aabs=*/2e-4, /*arel=*/2e-4})); @@ -357,9 +356,9 @@ ENTRY e { "block_level_fusion_config":{"output_tiles":[{"sizes":[16,32]}], "num_stages":4,"num_warps":8,"num_ctas":1}}} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTest, /*data_type=*/{}, HloOpcode::kDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTest, /*data_type=*/{}, HloOpcode::kDot)); EXPECT_THAT( legacy_triton::CanTritonHandleGEMM( *Cast(&ti.Instruction()), GetComputeCapability()) @@ -385,9 +384,9 @@ ENTRY e { "block_level_fusion_config":{"output_tiles":[{"sizes":[16,32]}], "num_stages":4,"num_warps":8,"num_ctas":1}}} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTest, /*data_type=*/{}, HloOpcode::kDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTest, /*data_type=*/{}, HloOpcode::kDot)); EXPECT_THAT( legacy_triton::CanTritonHandleGEMM( *Cast(&ti.Instruction()), GetComputeCapability()) @@ -415,9 +414,9 @@ ENTRY e { "block_level_fusion_config":{"output_tiles":[{"sizes":[16,32]}], "num_stages":4,"num_warps":8,"num_ctas":1}}} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTest, /*data_type=*/{}, HloOpcode::kDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTest, /*data_type=*/{}, HloOpcode::kDot)); EXPECT_THAT( legacy_triton::CanTritonHandleGEMM( *Cast(&ti.Instruction()), GetComputeCapability()) @@ -445,9 +444,9 @@ ENTRY e { "block_level_fusion_config":{"output_tiles":[{"sizes":[1,1,2,2]}], "num_stages":4,"num_warps":8,"num_ctas":1}}} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTest, /*data_type=*/{}, HloOpcode::kDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTest, /*data_type=*/{}, HloOpcode::kDot)); const se::DeviceDescription dev_info = TestGpuDeviceInfo::RTXA6000DeviceInfo(GetComputeCapability()); EXPECT_TRUE(legacy_triton::IsTritonSupportedInstruction( @@ -458,9 +457,9 @@ ENTRY e { .backend_config() ->fusion_backend_config() .block_level_fusion_config()); - TF_EXPECT_OK(TritonWrapper( - "test_fn", ti.TritonFusion(), GetComputeCapability(), dev_info, - block_level_parameters, target_triple_, data_layout_, mlir_context_)); + EXPECT_OK(TritonWrapper("test_fn", ti.TritonFusion(), GetComputeCapability(), + dev_info, block_level_parameters, target_triple_, + data_layout_, mlir_context_)); } TEST_F(SupportLegacyTest, @@ -482,9 +481,9 @@ ENTRY e { backend_config={"fusion_backend_config":{"kind":"__triton_nested_gemm_fusion", "block_level_fusion_config":{"output_tiles":[{"sizes":[1,1]}]}}} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTest, /*data_type=*/{}, HloOpcode::kDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTest, /*data_type=*/{}, HloOpcode::kDot)); EXPECT_THAT(legacy_triton::IsTritonSupportedInstruction( ti.Instruction(), GetComputeCapability()) .Explain(), diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/support_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/support_test.cc index 8ee5570c30a039..c162c3c2267d81 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/support_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/support_test.cc @@ -54,7 +54,6 @@ limitations under the License. #include "xla/stream_executor/cuda/cuda_compute_capability.h" #include "xla/stream_executor/device_description.h" #include "xla/stream_executor/rocm/rocm_compute_capability.h" -#include "xla/tsl/platform/statusor.h" #include "xla/xla.pb.h" #include "xla/xla_data.pb.h" #include "tsl/platform/protobuf.h" @@ -430,7 +429,7 @@ TEST_P(SupportTestWithTilingParam, IsTritonSupportedComputationSkipsRootTuple) { negate = f32[10] negate(abs) ROOT res = (f32[10], f32[10]) tuple(abs, negate) })"; - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(kHlo)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(kHlo)); EXPECT_TRUE(IsTritonSupportedComputation( *module->entry_computation(), se::CudaComputeCapability::Hopper())); } @@ -455,7 +454,7 @@ ENTRY triton_computation { parameter_0 = $0[1,16,4] parameter(0) ROOT bitcast_or_reshape = $0[64] $1(parameter_0) })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16}, cc); @@ -468,7 +467,7 @@ ENTRY triton_computation { parameter_0 = $0[1,1,1] parameter(0) ROOT bitcast_or_reshape = $0[] $1(parameter_0) })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{}, cc); @@ -492,7 +491,7 @@ ENTRY triton_computation { p1 = $0[] parameter(1) ROOT pad = $0[32, 16] $1(p0, p1), padding=0_28_0x0_12_0 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{4, 4}, cc); @@ -506,7 +505,7 @@ ENTRY triton_computation { p1 = $0[] parameter(1) ROOT pad = $0[7] $1(p0, p1), padding=0_0_1 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{4}, cc); @@ -520,7 +519,7 @@ ENTRY triton_computation { p1 = $0[] parameter(1) ROOT pad = $0[8] $1(p0, p1), padding=4_0_0 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{4}, cc); @@ -569,7 +568,7 @@ ENTRY triton_computation { bool f64_output = opcode == HloOpcode::kReal || opcode == HloOpcode::kImag || (opcode == HloOpcode::kAbs && primitive_util::IsComplexType(data_type)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction( f64_output ? kF64OutputTemplate @@ -652,7 +651,7 @@ ENTRY triton_computation { primitive_util::LowercasePrimitiveTypeName(data_type_in), primitive_util::LowercasePrimitiveTypeName(data_type_out)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction( hlo_text, data_type_in, // The type provided here is irrelevant. @@ -736,12 +735,11 @@ ENTRY triton_computation { ROOT compare = pred[11,63] $1(parameter_0, parameter_1), direction=GE })"; - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(opcode == HloOpcode::kCompare - ? kHloCompareTestTemplate - : kHloTestTemplate, - data_type, opcode)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + opcode == HloOpcode::kCompare + ? kHloCompareTestTemplate + : kHloTestTemplate, + data_type, opcode)); ExpectedFailMode fail_mode = ExpectedFailMode::kFail; if (cc.IsCuda()) { @@ -780,12 +778,11 @@ ENTRY triton_computation { ROOT compare = pred[] $1(parameter_0, parameter_1), direction=GE })"; - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(opcode == HloOpcode::kCompare - ? kHloCompareTestTemplate - : kHloTestTemplate, - data_type, opcode)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + opcode == HloOpcode::kCompare + ? kHloCompareTestTemplate + : kHloTestTemplate, + data_type, opcode)); ExpectedFailMode fail_mode = ExpectedFailMode::kFail; if (cc.IsCuda()) { @@ -854,9 +851,8 @@ ENTRY triton_computation { absl::Substitute(kHloTestTemplate, type, HloOpcodeString(opcode), opcode == HloOpcode::kSelect ? "pred" : type); - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(hlo_text, data_type, opcode)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + hlo_text, data_type, opcode)); bool skip_failure_branch_to_avoid_crash = false; if (cc.IsRocm()) { @@ -909,7 +905,7 @@ ENTRY triton_computation { dimensions={1}, to_apply=add })", init_value(data_type)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); bool crashes_on_failure = data_type == PrimitiveType::F8E4M3FN || @@ -937,9 +933,9 @@ ENTRY triton_computation { ROOT reduce = $0[3,125] reduce(parameter_0, constant_0), dimensions={2}, to_apply=add })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction(kHloTestTemplate, F32, - HloOpcode::kReduce)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction(kHloTestTemplate, F32, + HloOpcode::kReduce)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{3, 4}, DefaultDeviceForTesting()); } @@ -960,7 +956,7 @@ ENTRY triton_computation { dimensions={0,2,3}, to_apply=add })", init_value(data_type)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); if (!tiling) { @@ -988,7 +984,7 @@ ENTRY triton_computation { ROOT reduce = $$0[127] reduce(parameter_0, constant_0), dimensions={0}, to_apply=add })", init_value(data_type)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); @@ -1025,7 +1021,7 @@ ENTRY triton_computation { dimensions={1}, to_apply=add })", init_value(data_type)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTestMultipleOutputTiles(std::move(ti), @@ -1046,9 +1042,9 @@ ENTRY triton_computation { init = $0[] parameter(1) ROOT reduce = $0[125] reduce(parameter_0, init), dimensions={1}, to_apply=add })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction(kHloTestTemplate, F32, - HloOpcode::kReduce)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction(kHloTestTemplate, F32, + HloOpcode::kReduce)); EXPECT_TRUE(IsTritonSupportedInstruction(ti.Instruction(), cc)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2}, cc); } @@ -1069,7 +1065,7 @@ ENTRY triton_computation { dimensions={1}, to_apply=custom_call })", init_value(data_type)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1}, cc); @@ -1108,9 +1104,9 @@ ENTRY triton_computation { })", HloOpcodeString(opcode), init_value(data_type)); - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, HloOpcode::kReduce)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kReduce)); // TODO(b/361526623): Reduce the cases where emitter crashes. ExpectedFailMode fail_mode = ExpectedFailMode::kFail; @@ -1297,7 +1293,7 @@ ENTRY triton_computation { parameter_0 = $0[125,127,37] parameter(0) ROOT transpose = $0[127,37,125] $1(parameter_0), dimensions={1,2,0} })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); @@ -1344,7 +1340,7 @@ ENTRY triton_computation { p = $0[128,32] parameter(0) ROOT slice = $0[12,5] $1(p), slice={[116:128], [20:25]} })"); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); @@ -1358,7 +1354,7 @@ ENTRY triton_computation { p = f32[16,16,32] parameter(0) ROOT slice = f32[4,4,8] slice(p), slice={[2:10:2], [2:6], [3:11]} })"); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); @@ -1372,7 +1368,7 @@ ENTRY triton_computation { p = f32[16,16,32] parameter(0) ROOT slice = f32[4,4,8] slice(p), slice={[3:11:2], [2:6], [3:11]} })"); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); @@ -1420,9 +1416,9 @@ ENTRY triton_computation { p2 = $0[18,128,20] parameter(2) ROOT concatenate = $0[18,384,20] concatenate(p0, p1, p2), dimensions={1} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTestTemplate, F32, HloOpcode::kConcatenate)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction(kHloTestTemplate, F32, + HloOpcode::kConcatenate)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1, 64, 1}, cc); } @@ -1436,9 +1432,9 @@ ENTRY triton_computation { p1 = $0[18,128,20] parameter(1) ROOT concatenate = $0[18,191,20] concatenate(p0, p1), dimensions={1} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTestTemplate, F32, HloOpcode::kConcatenate)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction(kHloTestTemplate, F32, + HloOpcode::kConcatenate)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1, 64, 1}, cc); } @@ -1459,9 +1455,9 @@ ENTRY triton_computation { p2 = $0[128] parameter(2) ROOT result = $0[384] concatenate(p0, p1, p2), dimensions={0} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kConcatenate)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, + HloOpcode::kConcatenate)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{64}, cc); } @@ -1484,9 +1480,9 @@ ENTRY triton_computation { ROOT all-gather = $0[128,128] all-gather(input), replica_groups={{0,1}}, dimensions={1} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kAllGather)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kAllGather)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2, 2}, cc); } @@ -1499,10 +1495,9 @@ ENTRY triton_computation { replica_groups={{0,1}}, dimensions={0} ROOT all-gather-done = $0[256,32] all-gather-done(all-gather-start) })"; - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, - HloOpcode::kAllGatherStart)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, + HloOpcode::kAllGatherStart)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2, 2}, cc); } @@ -1521,9 +1516,9 @@ ENTRY triton_computation { ROOT all-reduce = $0[128,32] all-reduce(input), replica_groups={}, to_apply=apply_op })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kAllReduce)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kAllReduce)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2, 2}, cc); } @@ -1534,9 +1529,9 @@ ENTRY triton_computation { input = $0[128,32] parameter(0) ROOT a2a = ($0[128,32]) all-to-all(input), replica_groups={} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kAllToAll)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kAllToAll)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2, 2}, cc); } @@ -1548,7 +1543,7 @@ ENTRY triton_computation { ROOT collective-permute = $0[128,32] collective-permute(a), source_target_pairs={{1,0}, {0,1}, {2,2}} })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, HloOpcode::kCollectivePermute)); @@ -1569,11 +1564,11 @@ ENTRY triton_computation { ROOT done = $0[128,32] collective-permute-done(start) })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti_start, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, HloOpcode::kCollectivePermuteStart)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti_done, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, HloOpcode::kCollectivePermuteDone)); @@ -1596,9 +1591,9 @@ ENTRY triton_computation { ROOT result = $0[4] reduce-scatter(input), replica_groups={}, dimensions={0}, to_apply=apply_op })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kReduceScatter)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, + HloOpcode::kReduceScatter)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1}, cc); } @@ -1621,18 +1616,17 @@ ENTRY triton_computation { calls=async_computation ROOT async-done = $0[10] async-done(async-update), calls=async_computation })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti_start, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, HloOpcode::kAsyncStart)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti_update, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, HloOpcode::kAsyncUpdate)); - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti_done, - ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, - HloOpcode::kAsyncDone)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti_done, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kAsyncDone)); RunSupportTest(std::move(ti_start), /*output_tile_sizes=*/{1}, cc); RunSupportTest(std::move(ti_update), /*output_tile_sizes=*/{1}, cc); RunSupportTest(std::move(ti_done), /*output_tile_sizes=*/{1}, cc); @@ -1646,7 +1640,7 @@ ENTRY triton_computation { input = $0[128,32] parameter(0) ROOT result = $0[128,32] collective-broadcast(input), replica_groups={} })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, HloOpcode::kCollectiveBroadcast)); @@ -1660,9 +1654,9 @@ ENTRY triton_computation { ROOT replica_id = u32[] replica-id() })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kReplicaId)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kReplicaId)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{}, cc); } @@ -1673,9 +1667,9 @@ ENTRY triton_computation { ROOT partition_id = u32[] partition-id() })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kPartitionId)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, + HloOpcode::kPartitionId)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{}, cc); } @@ -1691,10 +1685,9 @@ ENTRY triton_computation { recv_sizes = s32[1] parameter(5) ROOT root = $0[128,32] ragged-all-to-all(input, output, input_offsets, send_sizes, output_offsets, recv_sizes), replica_groups={} })"; - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, - HloOpcode::kRaggedAllToAll)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, + HloOpcode::kRaggedAllToAll)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2, 2}, cc); } @@ -1740,9 +1733,9 @@ ENTRY triton_computation { input = $0[35,131] parameter(0) ROOT bcast = $0[3,35,131,12] broadcast(input), dimensions={1,2} })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kBroadcast)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kBroadcast)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2, 16, 32, 8}, cc); } @@ -1771,10 +1764,9 @@ ENTRY triton_computation { ROOT noop = s8[35,131] convert(input) })"; } - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(hlo_test_template, data_type, - HloOpcode::kParameter)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + hlo_test_template, data_type, + HloOpcode::kParameter)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 32}, cc); } @@ -1799,9 +1791,9 @@ ENTRY triton_computation { })", init_value(data_type)); - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kConstant)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kConstant)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1, 1}, cc); } @@ -1815,9 +1807,9 @@ ENTRY triton_computation { })", init_value(data_type)); - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kConstant)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kConstant)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2, 2}, cc); } @@ -1838,7 +1830,7 @@ TEST_P(IotaTest, Iota2D) { ENTRY triton_computation { ROOT input = $0[35,131] iota(), iota_dimension=0 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 32}, cc); @@ -1861,7 +1853,7 @@ ENTRY triton_computation { high = $0[] parameter(1) ROOT root = $0[33,77] rng(low, high), distribution=rng_uniform })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 32}, cc); @@ -1882,7 +1874,7 @@ ENTRY triton_computation { state = u64[2] parameter(0) ROOT root = (u64[2], $0[33,77]) rng-bit-generator(state), algorithm=rng_three_fry })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTestMultipleOutputTiles(std::move(ti), @@ -1902,7 +1894,7 @@ TEST_P(RngGetAndUpdateStateTest, RngGetAndUpdateState) { ENTRY triton_computation { ROOT root = u64[2]{0} rng-get-and-update-state(), delta=4096 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kRngGetAndUpdateState)); @@ -1933,7 +1925,7 @@ ENTRY triton_computation { ROOT root = c128[33,77] complex(real, imag) })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction( data_type == F32 ? kF32HloTestTemplate : kF64HloTestTemplate, @@ -1966,7 +1958,7 @@ ENTRY triton_computation { true_computation=true_branch, false_computation=false_branch })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1}, cc); @@ -1996,7 +1988,7 @@ ENTRY triton_computation { constant = s32[] constant(0) ROOT while = s32[] while(constant), condition=condition, body=body })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kWhile)); @@ -2023,7 +2015,7 @@ ENTRY triton_computation { operand = $0[10] parameter(0) ROOT call_op = $0[10] call(operand), to_apply=called_computation })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1}, cc); @@ -2049,7 +2041,7 @@ ENTRY triton_computation { ROOT bn_inf = $0[4,8,16,32] batch-norm-inference(operand, scale, offset, mean, variance), epsilon=0.001, feature_index=3 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1, 1, 4, 8}, cc); @@ -2073,7 +2065,7 @@ ENTRY triton_computation { ROOT bn_train = ($0[4,8,16,32], $0[32], $0[32]) batch-norm-training(operand, scale, offset), epsilon=0.001, feature_index=3 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTestMultipleOutputTiles( @@ -2099,7 +2091,7 @@ ENTRY triton_computation { ROOT bn_grad = ($0[4,8,16,32], $0[32], $0[32]) batch-norm-grad(operand, scale, mean, variance, grad_output), epsilon=0.001, feature_index=3 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTestMultipleOutputTiles( @@ -2120,7 +2112,7 @@ ENTRY triton_computation { operand = $0[] parameter(0) ROOT domain_op = $0[] domain(operand), domain={kind="sharding", entry={maximal device=0}, exit={maximal device=1}} })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{}, cc); @@ -2141,7 +2133,7 @@ ENTRY triton_computation { operand = s32[16, 32] parameter(0) ROOT get_dim_size = s32[] get-dimension-size(operand), dimensions={1} })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kGetDimensionSize)); @@ -2163,7 +2155,7 @@ ENTRY triton_computation { operand = $0[16,32] parameter(0) ROOT reverse_op = $0[16,32] reverse(operand), dimensions={0, 1} })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{4, 8}, cc); @@ -2233,10 +2225,9 @@ ENTRY triton_computation { hlo_text, primitive_util::LowercasePrimitiveTypeName(input_type), primitive_util::LowercasePrimitiveTypeName(result_type)); - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(hlo_text, PRIMITIVE_TYPE_INVALID, - HloOpcode::kDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + hlo_text, PRIMITIVE_TYPE_INVALID, HloOpcode::kDot)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 32}, cc, fail_mode); } @@ -2291,10 +2282,9 @@ triton_computation { hlo_text, primitive_util::LowercasePrimitiveTypeName(lhs_type), primitive_util::LowercasePrimitiveTypeName(rhs_type)); - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(hlo_text, PRIMITIVE_TYPE_INVALID, - HloOpcode::kDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + hlo_text, PRIMITIVE_TYPE_INVALID, HloOpcode::kDot)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 32}, cc); } @@ -2317,14 +2307,14 @@ triton_computation { } )"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti_rocm, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kDot)); RunSupportTest(std::move(ti_rocm), /*output_tile_sizes=*/{16, 32}, se::GpuComputeCapability(se::RocmComputeCapability("gfx942"))); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti_cuda, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kDot)); @@ -2343,7 +2333,7 @@ ENTRY triton_computation { backend_config={sizes:[64]} } )"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, F32, HloOpcode::kDot)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1, 16, 32}, @@ -2360,7 +2350,7 @@ ENTRY triton_computation { backend_config={sizes:[64]} } )"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, F32, HloOpcode::kDot)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1, 16, 1, 32}, @@ -2378,7 +2368,7 @@ ENTRY triton_computation { backend_config={"sizes":["64", "4"]} } )"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, F32, HloOpcode::kDot)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 32}, @@ -2397,7 +2387,7 @@ ENTRY triton_computation { backend_config={sizes:[64]} } )"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, F32, HloOpcode::kDot)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 32}, @@ -2416,7 +2406,7 @@ ENTRY triton_computation { backend_config={sizes:[64]} } )"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, F32, HloOpcode::kDot)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 32}, @@ -2470,7 +2460,7 @@ ENTRY triton_computation { if (absl::c_linear_search(std::vector{F8E5M2, F8E4M3FN, S8}, data_type)) { fail_mode = ExpectedFailMode::kFailOrCrash; } - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction( hlo_text, PrimitiveType::PRIMITIVE_TYPE_INVALID, HloOpcode::kDot)); @@ -2539,7 +2529,7 @@ ENTRY triton_computation { << "b/433240828: Triton fails on this combination in debug mode."; } - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(hlo_text, F32, HloOpcode::kDot)); ExpectedFailMode fail_mode = ExpectedFailMode::kFail; @@ -2584,9 +2574,9 @@ ENTRY triton_computation { backend_config={sizes:[16]} } )"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTestTemplate, type, HloOpcode::kScaledDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction(kHloTestTemplate, type, + HloOpcode::kScaledDot)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 16}, se::CudaComputeCapability::Hopper()); } @@ -2607,9 +2597,9 @@ ENTRY triton_computation { backend_config={sizes:[16]} } )"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTestTemplate, type, HloOpcode::kScaledDot)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction(kHloTestTemplate, type, + HloOpcode::kScaledDot)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{16, 16}, se::CudaComputeCapability::Hopper()); } @@ -2643,7 +2633,7 @@ ENTRY triton_computation { backend_config={sizes:[16]} } )"; - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(kHlo)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(kHlo)); const HloInstruction* dot = module->entry_computation()->root_instruction(); EXPECT_FALSE(IsTritonSupportedInstruction( *dot, se::GpuComputeCapability(se::CudaComputeCapability::Volta()))); @@ -2675,7 +2665,7 @@ ENTRY entry { ROOT fusion = bf16[16,64] fusion(p0, p1), kind=kCustom, calls=triton_computation } )"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(hlo_text, F32, HloOpcode::kFusion)); se::GpuComputeCapability cc = DefaultDeviceForTesting(); @@ -2742,10 +2732,9 @@ ENTRY triton_computation { output_tile_sizes = {1}; } - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(hlo_text, data_type_in, - HloOpcode::kBitcastConvert)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + hlo_text, data_type_in, HloOpcode::kBitcastConvert)); RunSupportTest(std::move(ti), output_tile_sizes, cc); } @@ -2795,9 +2784,9 @@ ENTRY triton_computation { output_tile_sizes = {1}; } - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction(hlo_text, data_type_in, - HloOpcode::kBitcast)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction(hlo_text, data_type_in, + HloOpcode::kBitcast)); RunSupportTest(std::move(ti), output_tile_sizes, cc); } @@ -2829,7 +2818,7 @@ ENTRY triton_computation { token0 = token[] after-all() ROOT add_dep = f32[10] add-dependency(param, token0) })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kAddDependency)); @@ -2852,7 +2841,7 @@ ENTRY triton_computation { token1 = token[] after-all() ROOT token2 = token[] after-all(token0, token1) })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kAfterAll)); @@ -2875,7 +2864,7 @@ ENTRY triton_computation { p1 = s32[5] parameter(1) ROOT tuple_op = (f32[10], s32[5]) tuple(p0, p1) })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kTuple)); @@ -2898,7 +2887,7 @@ ENTRY triton_computation { tuple_op = (f32[10], s32[5]) parameter(0) ROOT gte = f32[10] get-tuple-element(tuple_op), index=0 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kGetTupleElement)); @@ -2920,7 +2909,7 @@ ENTRY triton_computation { parameter = f32[10] parameter(0) ROOT custom_call_op = f32[10] custom-call(parameter), custom_call_target="SomeTarget" })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kCustomCall)); @@ -2955,9 +2944,9 @@ ENTRY triton_computation { })", primitive_util::LowercasePrimitiveTypeName(data_type), lower); - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kCholesky)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kCholesky)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2, 2}, cc); } @@ -3031,10 +3020,9 @@ ENTRY triton_computation { lower ? "true" : "false", unit_diagonal ? "true" : "false", TriangularSolveOptions::Transpose_Name(transpose_a)); - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(hlo_text, data_type, - HloOpcode::kTriangularSolve)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + hlo_text, data_type, HloOpcode::kTriangularSolve)); RunSupportTest(std::move(ti), {1, 2, 1}, cc); } @@ -3053,10 +3041,9 @@ ENTRY triton_computation { lower ? "true" : "false", unit_diagonal ? "true" : "false", TriangularSolveOptions::Transpose_Name(transpose_a)); - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti, - ParseTemplateAndGetInstruction(hlo_text, data_type, - HloOpcode::kTriangularSolve)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + hlo_text, data_type, HloOpcode::kTriangularSolve)); RunSupportTest(std::move(ti), {1, 1, 2}, cc); } @@ -3071,7 +3058,7 @@ ENTRY triton_computation { ROOT fft_op = $0[16,16] fft(parameter), fft_type=FFT, fft_length={16} })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(hlo_text, data_type, HloOpcode::kFft)); @@ -3087,7 +3074,7 @@ ENTRY triton_computation { ROOT fft_op = $0[16,16] fft(parameter), fft_type=IFFT, fft_length={16} })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(hlo_text, data_type, HloOpcode::kFft)); @@ -3111,7 +3098,7 @@ ENTRY triton_computation { })", real_data_type_str, complex_data_type_str); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(hlo_text, data_type, HloOpcode::kFft)); @@ -3135,7 +3122,7 @@ ENTRY triton_computation { })", complex_data_type_str, real_data_type_str); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(hlo_text, data_type, HloOpcode::kFft)); @@ -3163,16 +3150,14 @@ ENTRY triton_computation { ROOT cp_done = $0[10,10,10] copy-done(cp_start) })"; - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti_start, - ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, - HloOpcode::kCopyStart)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti_start, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kCopyStart)); RunSupportTest(std::move(ti_start), /*output_tile_sizes=*/{1, 1, 1}, cc); - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti_done, - ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, - HloOpcode::kCopyDone)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti_done, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kCopyDone)); RunSupportTest(std::move(ti_done), /*output_tile_sizes=*/{1, 1, 1}, cc); } constexpr std::array kTestedOpsCopy = {HloOpcode::kCopyStart, @@ -3194,9 +3179,9 @@ ENTRY triton_computation { token0 = token[] after-all() ROOT infeed_op = ($0[10], token[]) infeed(token0) })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, HloOpcode::kInfeed)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kInfeed)); RunSupportTestMultipleOutputTiles(std::move(ti), /*output_tile_sizes=*/{{1}, {}}, cc); } @@ -3218,9 +3203,9 @@ ENTRY triton_computation { token0 = token[] after-all() ROOT outfeed_op = token[] outfeed(data, token0) })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kOutfeed)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kOutfeed)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{}, cc); } @@ -3240,7 +3225,7 @@ ENTRY triton_computation { parameter = $0[10, 20] parameter(0) ROOT map_op = $0[10, 20] map(parameter), dimensions={0, 1}, to_apply=map_computation })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{4, 8}, cc); @@ -3274,7 +3259,7 @@ ENTRY triton_computation { operand = $0[10,20,30] parameter(0) ROOT sort_op = $0[10,20,30] sort(operand), dimensions={2}, is_stable=true, to_apply=compare })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{2, 4, 8}, cc); @@ -3297,7 +3282,7 @@ ENTRY triton_computation { values = s32[10,20] parameter(1) ROOT sort_op = ($0[10,20], s32[10,20]) sort(keys, values), dimensions={1}, is_stable=true, to_apply=compare })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, opcode)); RunSupportTestMultipleOutputTiles(std::move(ti), @@ -3320,15 +3305,14 @@ TEST_P(RecvOpsTest, RecvAndRecvDone) { recv_done_op = ($0[10,20], token[]) recv-done(recv_op), channel_id=15 ROOT result = $0[10,20] get-tuple-element(recv_done_op), index=0 })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti_recv, - ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, HloOpcode::kRecv)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti_recv, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kRecv)); RunSupportTest(std::move(ti_recv), /*output_tile_sizes=*/{1, 1}, cc); - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti_recv_done, - ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, - HloOpcode::kRecvDone)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti_recv_done, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kRecvDone)); RunSupportTest(std::move(ti_recv_done), /*output_tile_sizes=*/{1, 1}, cc); } @@ -3353,15 +3337,14 @@ ENTRY triton_computation { ROOT send_done_op = token[] send-done(send_op), channel_id=77 })"; - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti_send, - ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, HloOpcode::kSend)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti_send, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kSend)); RunSupportTest(std::move(ti_send), /*output_tile_sizes=*/{}, cc); - TF_ASSERT_OK_AND_ASSIGN( - TestedInstruction ti_send_done, - ParseTemplateAndGetInstruction(kHloTestTemplate, data_type, - HloOpcode::kSendDone)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti_send_done, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kSendDone)); RunSupportTest(std::move(ti_send_done), /*output_tile_sizes=*/{}, cc); } @@ -3405,7 +3388,7 @@ ENTRY triton_computation { primitive_util::LowercasePrimitiveTypeName(random_type), primitive_util::LowercasePrimitiveTypeName(new_element_type)); - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(hlo_text, PRIMITIVE_TYPE_INVALID, HloOpcode::kStochasticConvert)); @@ -3460,9 +3443,9 @@ ENTRY triton_computation { ROOT topk_op = ($$0[11,33,10], s32[11,33,10]) topk(operand), k=10, largest=$0 })", largest); - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, - ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, HloOpcode::kTopK)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, + ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, HloOpcode::kTopK)); RunSupportTestMultipleOutputTiles( std::move(ti), /*output_tile_sizes=*/{{2, 2, 1}, {2, 2, 1}}, cc); @@ -3501,9 +3484,9 @@ ENTRY triton_computation { })", primitive_util::LowercasePrimitiveTypeName(data_type), PrecisionToString(input_precision), PrecisionToString(kernel_precision)); - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kConvolution)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, + HloOpcode::kConvolution)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1, 2, 2, 1}, cc); } @@ -3521,9 +3504,9 @@ ENTRY triton_computation { })", primitive_util::LowercasePrimitiveTypeName(data_type), PrecisionToString(input_precision), PrecisionToString(kernel_precision)); - TF_ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( - kHloTestTemplate, data_type, - HloOpcode::kConvolution)); + ASSERT_OK_AND_ASSIGN(TestedInstruction ti, ParseTemplateAndGetInstruction( + kHloTestTemplate, data_type, + HloOpcode::kConvolution)); RunSupportTest(std::move(ti), /*output_tile_sizes=*/{1, 1, 2, 2}, cc); } @@ -3548,7 +3531,7 @@ ENTRY triton_computation { ROOT conv = f16[1,2,2,3] convolution(input, kernel), window={size=3x3 stride=2x2}, dim_labels=b01f_01io->b01f })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); @@ -3565,7 +3548,7 @@ ENTRY triton_computation { ROOT conv = f16[1,1,2,3] convolution(input, kernel), window={size=3x3 rhs_dilate=2x2}, dim_labels=b01f_01io->b01f })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); @@ -3584,7 +3567,7 @@ ENTRY triton_computation { feature_group_count=2 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); @@ -3602,7 +3585,7 @@ ENTRY triton_computation { window={size=3x3 pad=1_1x1_1}, dim_labels=b01f_01io->b01f, batch_group_count=2 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); @@ -3619,7 +3602,7 @@ ENTRY triton_computation { ROOT conv = f16[1,7,9,3] convolution(input, kernel), window={size=3x3 lhs_dilate=2x2}, dim_labels=b01f_01io->b01f })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); @@ -3636,7 +3619,7 @@ ENTRY triton_computation { ROOT conv = f16[1,5,7,3] convolution(input, kernel), window={size=3x3 pad=1_1x1_2}, dim_labels=b01f_01io->b01f })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); @@ -3654,7 +3637,7 @@ ENTRY triton_computation { window={size=3x3 pad=1_1x1_1}, dim_labels=b01f_01io->b01f, feature_group_count=2 })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); @@ -3673,7 +3656,7 @@ ENTRY triton_computation { ROOT conv = f16[1,2,2,3] convolution(input, kernel), window={size=3x3 pad=1_1x1_1}, dim_labels=b01f_01io->b01f })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); @@ -3689,7 +3672,7 @@ ENTRY triton_computation { ROOT conv = f16[1,5,6,3] convolution(input, kernel), window={size=3x3 pad=2_0x0_2}, dim_labels=b01f_01io->b01f })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); @@ -3705,7 +3688,7 @@ ENTRY triton_computation { ROOT conv = f16[1,5,6,3] convolution(input, kernel), window={size=2x2 pad=1_0x0_1}, dim_labels=b01f_01io->b01f })"; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TestedInstruction ti, ParseTemplateAndGetInstruction(kHloTestTemplate, PRIMITIVE_TYPE_INVALID, HloOpcode::kConvolution)); diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/tests/BUILD b/third_party/xla/xla/backends/gpu/codegen/triton/tests/BUILD index 1532db909a751e..21e32f5e9816ba 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/tests/BUILD +++ b/third_party/xla/xla/backends/gpu/codegen/triton/tests/BUILD @@ -201,10 +201,8 @@ xla_test( "//xla/tests:hlo_interpreter_reference_mixin", "//xla/tests:test_utils", "//xla/tests:xla_internal_test_main", # fixdeps: keep - "//xla/tsl/lib/core:status_test_util", "//xla/tsl/platform:env", "//xla/tsl/platform:errors", - "//xla/tsl/platform:statusor", "//xla/tsl/platform:test", "@com_google_absl//absl/algorithm:container", "@com_google_absl//absl/log", @@ -213,6 +211,7 @@ xla_test( "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_macros", "@com_google_absl//absl/status:status_matchers", + "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", "@com_google_absl//absl/types:span", "@com_google_googletest//:gtest", @@ -263,7 +262,6 @@ xla_test( "//xla/tests:hlo_interpreter_reference_mixin", "//xla/tests:test_utils", "//xla/tests:xla_internal_test_main", # fixdeps: keep - "//xla/tsl/platform:statusor", "@com_google_absl//absl/algorithm:container", "@com_google_absl//absl/log", "@com_google_absl//absl/log:check", @@ -271,6 +269,7 @@ xla_test( "@com_google_absl//absl/status", "@com_google_absl//absl/status:status_macros", "@com_google_absl//absl/status:status_matchers", + "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", "@com_google_absl//absl/types:span", "@com_google_googletest//:gtest", 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 22f2360d567de0..a6cfab6a822544 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 @@ -36,6 +36,7 @@ limitations under the License. #include "absl/status/status.h" #include "absl/status/status_macros.h" #include "absl/status/status_matchers.h" +#include "absl/status/statusor.h" #include "absl/strings/match.h" #include "absl/strings/str_cat.h" #include "absl/strings/str_join.h" @@ -84,10 +85,8 @@ limitations under the License. #include "xla/stream_executor/rocm/rocm_compute_capability.h" #include "xla/tests/hlo_interpreter_reference_mixin.h" #include "xla/tests/test_utils.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/tsl/platform/env.h" #include "xla/tsl/platform/errors.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/platform/test.h" #include "xla/types.h" #include "xla/util.h" @@ -542,7 +541,7 @@ CHECK: return CHECK: } )")); - TF_EXPECT_OK(LowerXTileIrToTritonAndFileCheck( + EXPECT_OK(LowerXTileIrToTritonAndFileCheck( xtile_module_and_hlo_module.first.get(), R"( CHECK: xtile.entry_func @xtile_dialect_fn(%[[P0:.*]]: {{.*}}, %[[P1:.*]]: {{.*}}, %[[PID:[^:]*]]: index{{( \{xla.range = \[0 : index, 124 : index\]\})?}}) CHECK-DAG: %[[C_0:.*]] = arith.constant 0 : index @@ -619,7 +618,7 @@ CHECK-DAG: xtile.insert {{.*}} into %[[P2]] CHECK-SAME: [%[[TID]], %{{.*}}] [1, 128] [1, 1] : tensor<1x128xf32> )")); - TF_EXPECT_OK(LowerXTileIrToTritonAndFileCheck( + EXPECT_OK(LowerXTileIrToTritonAndFileCheck( xtile_module_and_hlo_module.first.get(), R"( CHECK: xtile.entry_func @xtile_dialect_fn( CHECK-SAME: %[[P0:[A-Za-z0-9_]*]]: memref<125x127xf32> @@ -882,7 +881,7 @@ CHECK: stablehlo.multiply {{.*}} tensor<1xf32> CHECK: xtile.insert {{.*}} : tensor<1xf32> )")); - TF_EXPECT_OK(LowerXTileIrToTritonAndFileCheck( + EXPECT_OK(LowerXTileIrToTritonAndFileCheck( xtile_module_and_hlo_module.first.get(), R"( CHECK: xtile.entry_func @xtile_dialect_fn(%[[P0:[A-Za-z0-9_]*]]: memref<125x127xf32> CHECK-SAME: %[[P1:[A-Za-z0-9_]*]]: memref<125xf32> @@ -1783,7 +1782,7 @@ ENTRY entry { })"; // Check that the IR attribute is set correctly. - TF_EXPECT_OK(CreateTritonIrFromHloTextAndFileCheck(hlo_text, "fdot", R"( + EXPECT_OK(CreateTritonIrFromHloTextAndFileCheck(hlo_text, "fdot", R"( // CHECK: scf.for // CHECK: scf.yield // CHECK-NEXT: tt.warp_specialize @@ -2251,7 +2250,7 @@ TEST_P(TritonEmitterTestWithTilingParam, RocmWarpSizeIsSetCorrectly) { "test_fn", *triton_fusion, se::GpuComputeCapability{se::RocmComputeCapability("gfx942")}, dev_info, block_level_parameters, target_triple, data_layout, mlir_context)); - TF_EXPECT_OK(tsl::Env::Default()->GetMatchingPaths( + EXPECT_OK(tsl::Env::Default()->GetMatchingPaths( tsl::io::JoinPath(output_directory, "*.triton-to-llvm.txt"), &paths)); EXPECT_EQ(paths.size(), 1); ASSERT_OK( @@ -2272,7 +2271,7 @@ TEST_P(TritonEmitterTestWithTilingParam, RocmWarpSizeIsSetCorrectly) { se::GpuComputeCapability{se::RocmComputeCapability("gfx1100")}, dev_info_n, block_level_parameters, target_triple, data_layout, mlir_context)); - TF_EXPECT_OK(tsl::Env::Default()->GetMatchingPaths( + EXPECT_OK(tsl::Env::Default()->GetMatchingPaths( tsl::io::JoinPath(output_directory, "*.triton-to-llvm.txt"), &paths)); EXPECT_EQ(paths.size(), 1); ASSERT_OK( diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/tests/scaled_dot_device_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/tests/scaled_dot_device_test.cc index d8a21a19a5f94a..b6f86813147cb8 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/tests/scaled_dot_device_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/tests/scaled_dot_device_test.cc @@ -35,6 +35,7 @@ limitations under the License. #include "absl/status/status.h" #include "absl/status/status_macros.h" #include "absl/status/status_matchers.h" +#include "absl/status/statusor.h" #include "absl/strings/match.h" #include "absl/strings/str_cat.h" #include "absl/strings/str_join.h" @@ -72,7 +73,6 @@ limitations under the License. #include "xla/stream_executor/device_description.h" #include "xla/tests/hlo_interpreter_reference_mixin.h" #include "xla/tests/test_utils.h" -#include "xla/tsl/platform/statusor.h" #include "xla/types.h" #include "xla/xla.pb.h" #include "xla/xla_data.pb.h" @@ -905,7 +905,7 @@ ENTRY e { "num_warps":"4","num_ctas":"1","num_stages":"1"}}} } )hlo"; - TF_ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(kHloText)); + ASSERT_OK_AND_ASSIGN(auto module, ParseAndReturnVerifiedModule(kHloText)); HloComputation* scaled_dot_computation = GetFirstComputationWithInstruction(*module, HloOpcode::kScaledDot); EXPECT_THAT(CreateTritonIrAndFileCheckForDot(*scaled_dot_computation, R"( diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/tests/triton_test_correctness.cc b/third_party/xla/xla/backends/gpu/codegen/triton/tests/triton_test_correctness.cc index 67d0c14fca2058..9da57e05145de6 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/tests/triton_test_correctness.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/tests/triton_test_correctness.cc @@ -17,6 +17,7 @@ limitations under the License. #include #include +#include #include #include "absl/log/log.h" #include "xla/debug_options_flags.h" @@ -25,7 +26,6 @@ limitations under the License. #include "xla/tests/hlo_pjrt_interpreter_reference_mixin.h" #include "xla/tests/hlo_pjrt_test_base.h" #include "xla/tools/hlo_module_loader.h" -#include "xla/tsl/platform/statusor.h" #include "xla/tsl/util/command_line_flags.h" namespace xla::gpu { @@ -39,8 +39,8 @@ float rel_error_bound = 0.0; using CorrectnessTest = HloInterpreterReferenceMixin; TEST_F(CorrectnessTest, RunAndCompare) { - TF_ASSERT_OK_AND_ASSIGN(std::unique_ptr module, - LoadModuleFromFile(input_file)); + ASSERT_OK_AND_ASSIGN(std::unique_ptr module, + LoadModuleFromFile(input_file)); EXPECT_TRUE(RunAndCompareNoHloPasses( std::move(module), ErrorSpec{abs_error_bound, rel_error_bound})); } diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/tma_utils_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/tma_utils_test.cc index 8af9a765b90034..5a20de4d3fdecc 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/tma_utils_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/tma_utils_test.cc @@ -26,7 +26,6 @@ limitations under the License. #include "mlir/IR/MLIRContext.h" #include "xla/backends/gpu/codegen/triton/ir/triton_xla_ops.h" #include "xla/stream_executor/gpu/tma_metadata.h" -#include "xla/tsl/platform/statusor.h" namespace xla::gpu { namespace { @@ -47,7 +46,7 @@ TEST(CreateTmaDescriptorTest, Valid2DInputReturnCorrectDescriptor) { llvm::SmallVector layout = {1, 0}; int element_byte_size = 4; SwizzleMode swizzle_mode = SwizzleMode::k128b; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TmaDescriptor tma_desc, CreateTmaDescriptor(global_shape, tile_shape, tile_strides, layout, element_byte_size, swizzle_mode)); @@ -72,7 +71,7 @@ TEST(CreateTmaDescriptorTest, Valid1DInputReturnCorrectDescriptor) { llvm::SmallVector layout = {0}; int element_byte_size = 4; SwizzleMode swizzle_mode = SwizzleMode::kNone; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TmaDescriptor tma_desc, CreateTmaDescriptor(global_shape, tile_shape, tile_strides, layout, element_byte_size, swizzle_mode)); @@ -97,7 +96,7 @@ TEST(CreateTmaDescriptorTest, Valid5DInputReturnCorrectDescriptor) { llvm::SmallVector layout = {4, 3, 2, 1, 0}; int element_byte_size = 2; SwizzleMode swizzle_mode = SwizzleMode::kNone; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TmaDescriptor tma_desc, CreateTmaDescriptor(global_shape, tile_shape, tile_strides, layout, element_byte_size, swizzle_mode)); @@ -172,7 +171,7 @@ TEST(CreateTmaDescriptorTest, NonUnitTileStridesAreCorrectlyHandled) { llvm::SmallVector layout = {1, 0}; int element_byte_size = 4; SwizzleMode swizzle_mode = SwizzleMode::k128b; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TmaDescriptor tma_desc, CreateTmaDescriptor(global_shape, tile_shape, tile_strides, layout, element_byte_size, swizzle_mode)); @@ -194,7 +193,7 @@ TEST(CreateTmaDescriptorTest, BoxDimsAreAdjustedForSwizzleMode) { // 128B swizzle mode. SwizzleMode swizzle_mode = SwizzleMode::k128b; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( TmaDescriptor tma_desc, CreateTmaDescriptor(global_shape, tile_shape, tile_strides, layout, element_byte_size, swizzle_mode)); @@ -202,14 +201,14 @@ TEST(CreateTmaDescriptorTest, BoxDimsAreAdjustedForSwizzleMode) { // 64B swizzle mode. swizzle_mode = SwizzleMode::k64b; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( tma_desc, CreateTmaDescriptor(global_shape, tile_shape, tile_strides, layout, element_byte_size, swizzle_mode)); EXPECT_EQ(tma_desc.box_dims()[0], 16); // 32B swizzle mode. swizzle_mode = SwizzleMode::k32b; - TF_ASSERT_OK_AND_ASSIGN( + ASSERT_OK_AND_ASSIGN( tma_desc, CreateTmaDescriptor(global_shape, tile_shape, tile_strides, layout, element_byte_size, swizzle_mode)); EXPECT_EQ(tma_desc.box_dims()[0], 8); diff --git a/third_party/xla/xla/backends/gpu/codegen/triton/triton_gemm_fusion_test.cc b/third_party/xla/xla/backends/gpu/codegen/triton/triton_gemm_fusion_test.cc index 3186968ae3c3b3..82e44682615713 100644 --- a/third_party/xla/xla/backends/gpu/codegen/triton/triton_gemm_fusion_test.cc +++ b/third_party/xla/xla/backends/gpu/codegen/triton/triton_gemm_fusion_test.cc @@ -57,7 +57,6 @@ limitations under the License. #include "xla/stream_executor/cuda/cuda_compute_capability.h" #include "xla/stream_executor/device_description.h" #include "xla/tests/hlo_interpreter_reference_mixin.h" -#include "xla/tsl/lib/core/status_test_util.h" #include "xla/tsl/platform/env.h" #include "xla/tsl/platform/errors.h" #include "xla/tsl/platform/test.h" @@ -914,11 +913,11 @@ ENTRY entry { const HloFusionInstruction* fusion2 = Cast( module1_and_metadata.computation->FusionInstruction()); - TF_EXPECT_OK(TritonWrapper("test_fn", *fusion2, se::GpuComputeCapability{cc}, - device_info, - module2_and_metadata.block_level_parameters, - target_triple, data_layout, mlir_context_) - .status()); + EXPECT_OK(TritonWrapper("test_fn", *fusion2, se::GpuComputeCapability{cc}, + device_info, + module2_and_metadata.block_level_parameters, + target_triple, data_layout, mlir_context_) + .status()); } // TODO(b/393299275): this test may have some value while Triton tiling From a3165c23b74c125ccab52dd1947c89b54358c086 Mon Sep 17 00:00:00 2001 From: "A. Unique TensorFlower" Date: Wed, 30 Sep 2026 05:09:17 -0700 Subject: [PATCH 58/58] Automated Code Change PiperOrigin-RevId: 990918017 --- tensorflow/core/lib/db/sqlite_test.cc | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tensorflow/core/lib/db/sqlite_test.cc b/tensorflow/core/lib/db/sqlite_test.cc index 99e50cde01f2fc..1c21a0de890f34 100644 --- a/tensorflow/core/lib/db/sqlite_test.cc +++ b/tensorflow/core/lib/db/sqlite_test.cc @@ -170,7 +170,7 @@ TEST_F(SqliteTest, UnsafeColumn) { stmt = db_->PrepareOrDie("SELECT b FROM T ORDER BY a"); TF_ASSERT_OK(stmt.Step(&is_done_)); absl::string_view p = stmt.ColumnStringUnsafe(0); - EXPECT_EQ('h', *p.data()); + EXPECT_EQ('h', p[0]); TF_ASSERT_OK(stmt.Step(&is_done_)); // This will actually happen, but it's not safe to test this behavior. // EXPECT_EQ('t', *p.data());