diff --git a/ci/official/utilities/extract_resultstore_links.py b/ci/official/utilities/extract_resultstore_links.py index 2bd96e1811c171..b24b4df7129277 100644 --- a/ci/official/utilities/extract_resultstore_links.py +++ b/ci/official/utilities/extract_resultstore_links.py @@ -114,8 +114,7 @@ def parse_log(file_path: str, else: 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_fragment = '\n'.join(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) diff --git a/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc b/tensorflow/core/kernels/conv_grad_filter_ops_3d.cc index f657da6337b410..640ddf447ab606 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 @@ -880,6 +881,14 @@ 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..4f0e24e6fe2ed1 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 @@ -399,6 +400,14 @@ 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, diff --git a/tensorflow/core/kernels/conv_grad_input_ops.cc b/tensorflow/core/kernels/conv_grad_input_ops.cc index e56826e7ddf580..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" @@ -291,6 +292,11 @@ void LaunchConv2DBackpropInputOpGpuImpl( : TensorShape({filter.dim_size(3), filter.dim_size(0), filter.dim_size(1), filter.dim_size(2)}); + 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()( diff --git a/tensorflow/core/kernels/conv_grad_input_ops_3d.cc b/tensorflow/core/kernels/conv_grad_input_ops_3d.cc index dc991c43c2b726..beaf66db1fcc28 100644 --- a/tensorflow/core/kernels/conv_grad_input_ops_3d.cc +++ b/tensorflow/core/kernels/conv_grad_input_ops_3d.cc @@ -17,6 +17,7 @@ limitations under the License. #define EIGEN_USE_THREADS #include +#include #include #include #include @@ -873,6 +874,12 @@ 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(context, + 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)); diff --git a/tensorflow/core/kernels/conv_ops_fused_impl.h b/tensorflow/core/kernels/conv_ops_fused_impl.h index 53e6804f015308..01bc404c3b7140 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 @@ -558,8 +559,14 @@ struct LaunchFusedConv2DOp { : TensorShape({filter.dim_size(3), filter.dim_size(0), filter.dim_size(1), filter.dim_size(2)}); + 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, To32Bit(filter.tensor()), diff --git a/tensorflow/core/kernels/conv_ops_impl.h b/tensorflow/core/kernels/conv_ops_impl.h index 5c84dd9bac12ca..03bd0815cea352 100644 --- a/tensorflow/core/kernels/conv_ops_impl.h +++ b/tensorflow/core/kernels/conv_ops_impl.h @@ -1133,6 +1133,12 @@ void LaunchConvOpImpl(OpKernelContext* context, bool cudnn_use_autotune, } } TensorShape dst_shape(dst_shape_vec); + OP_REQUIRES(context, + 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)); 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, diff --git a/tensorflow/core/kernels/image/crop_and_resize_op.cc b/tensorflow/core/kernels/image/crop_and_resize_op.cc index ff09def018cddb..4f9063471dd510 100644 --- a/tensorflow/core/kernels/image/crop_and_resize_op.cc +++ b/tensorflow/core/kernels/image/crop_and_resize_op.cc @@ -55,24 +55,21 @@ 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())); + 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, got ", box_index.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 (box_index.dim_size(0) != *num_boxes) { return absl::InvalidArgumentError("box_index has incompatible shape"); } diff --git a/tensorflow/core/kernels/split_v_op.cc b/tensorflow/core/kernels/split_v_op.cc index bd89be3b2dff37..d14e685c5fd679 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 @@ -100,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) { @@ -128,6 +137,24 @@ 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. + OP_REQUIRES(context, + determined_size <= std::numeric_limits::max() - size, + errors::InvalidArgument( + "Sum of size_splits overflows the index type at index ", + d, ".")); determined_size += size; } } @@ -147,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/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()); 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..1a099db8dd502c 100644 --- a/tensorflow/python/kernel_tests/array_ops/split_op_test.py +++ b/tensorflow/python/kernel_tests/array_ops/split_op_test.py @@ -120,6 +120,63 @@ def testExplicitNum(self): self.assertAllEqual(r[1], value[2:4]) 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. + i64_max = (1 << 63) - 1 + i32_max = (1 << 31) - 1 + 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, "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( + constant_op.constant([], dtype=dtypes.float32), + constant_op.constant([input_size, 0], dtype=dtypes.int64), + ) + with self.assertRaisesRegex( + (ValueError, errors_impl.InvalidArgumentError), message + ): + self.evaluate( + array_ops.split( + value, constant_op.constant(size_splits, dtype=dtype), axis=0 + ) + ) + + @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) diff --git a/tensorflow/python/ops/image_grad_test_base.py b/tensorflow/python/ops/image_grad_test_base.py index 699c4bbee3a964..3ef0dbd3c9d60c 100644 --- a/tensorflow/python/ops/image_grad_test_base.py +++ b/tensorflow/python/ops/image_grad_test_base.py @@ -448,6 +448,93 @@ 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 = 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) + 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, + ) + ) + # 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=image, + boxes=array_ops.zeros([0], dtype=dtypes.float32), + 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' + ): + 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, + boxes=valid_boxes, + box_ind=valid_box_ind, + 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. diff --git a/tensorflow/python/ops/numpy_ops/np_array_ops.py b/tensorflow/python/ops/numpy_ops/np_array_ops.py index 0455ea9bd48a16..da95dbdb107968 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops.py @@ -987,8 +987,37 @@ 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. 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)}.' + ) + normalized_mask = 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'array of rank {maybe_rank}.' + ) + bit = 1 << normalized + if normalized_mask & bit: + raise ValueError('repeated axis in transpose') + normalized_mask |= bit + 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) @@ -1631,6 +1660,41 @@ 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 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. + 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) + ): + # 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 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 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 085ee6fe95273b..0d47b3b2e6e2b9 100644 --- a/tensorflow/python/ops/numpy_ops/np_array_ops_test.py +++ b/tensorflow/python/ops/numpy_ops/np_array_ops_test.py @@ -1213,6 +1213,33 @@ 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]) + # 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: diff --git a/tensorflow/python/ops/numpy_ops/tests/np_test.py b/tensorflow/python/ops/numpy_ops/tests/np_test.py index d7116bce75568d..3843d73c936d39 100644 --- a/tensorflow/python/ops/numpy_ops/tests/np_test.py +++ b/tensorflow/python/ops/numpy_ops/tests/np_test.py @@ -2075,6 +2075,50 @@ 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,)) + 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)) + # 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=(onp.uint32(0), onp.uint32(0))) + with self.assertRaisesRegex(ValueError, "must be different"): + 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=(onp.uint32(0), onp.uint32(1))), + onp.rot90(onp.ones((2, 3)), axes=(0, 1)), + check_dtypes=False, + ) + # TODO(mattjj): test infix operator overrides def testRavel(self): diff --git a/tensorflow/tf_framework_version_script.lds b/tensorflow/tf_framework_version_script.lds index 99ed72972e3ba6..4a4956bb56fefe 100644 --- a/tensorflow/tf_framework_version_script.lds +++ b/tensorflow/tf_framework_version_script.lds @@ -1,4 +1,20 @@ tensorflow { global: *; -}; \ No newline at end of file + + # 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. + local: + LLVM*; +}; diff --git a/tensorflow/tf_private_symbols.lds b/tensorflow/tf_private_symbols.lds index 319b40bc72b66a..264a68e0d7f454 100644 --- a/tensorflow/tf_private_symbols.lds +++ b/tensorflow/tf_private_symbols.lds @@ -6,4 +6,5 @@ _jzero_far _jcopy_* _jsimd_* _hwloc_* +_LLVM* *mlir* 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 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 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/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/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/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/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_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/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/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/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/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 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/runtime/BUILD b/third_party/xla/xla/backends/gpu/runtime/BUILD index f5c3acd8945ea8..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", ], ) @@ -4339,6 +4340,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 +4348,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", @@ -5438,8 +5441,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 +5496,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/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/runtime/annotation.cc b/third_party/xla/xla/backends/gpu/runtime/annotation.cc index af9ff4e116e475..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 @@ -548,12 +558,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 +591,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,21 +604,26 @@ 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_)) { + // 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_; return; } InstructionAnnotationMetadata metadata = GetInstructionAnnotationMetadata(inst); + xprof_name_str_ = MakeInstructionName(nvtx_name_str_, metadata); payload_ = Detailed{ std::move(std::get(payload_)), @@ -783,6 +791,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 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 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 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 c7fd9b12c2af57..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", @@ -1708,7 +1696,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", @@ -1908,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_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/backends/gpu/tests/all_reduce_e2e_test.cc b/third_party/xla/xla/backends/gpu/tests/all_reduce_e2e_test.cc index 97780efd450c51..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}); @@ -1383,5 +1382,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/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 f59c95a6dc0b29..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) { @@ -194,6 +192,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(); @@ -351,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) { @@ -470,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) { @@ -522,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_) { @@ -581,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_) { @@ -633,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_) { @@ -682,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()) { @@ -729,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()) { @@ -877,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()) { @@ -919,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); @@ -954,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()) { @@ -1001,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()) { @@ -1074,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; @@ -1095,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); @@ -1293,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) { @@ -1338,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) { @@ -1382,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. @@ -1426,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) { @@ -1478,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]); @@ -1506,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) { @@ -1567,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_; @@ -1629,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]); @@ -1661,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]); @@ -1692,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]); @@ -1726,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}); @@ -1868,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] @@ -1935,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] @@ -2057,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; @@ -2103,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. @@ -2151,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 = @@ -2189,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=*/{}, @@ -2239,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)); @@ -2265,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); @@ -2314,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()); @@ -2906,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); @@ -2929,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; @@ -3141,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 = @@ -3180,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; @@ -3278,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); } @@ -3393,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}; @@ -3446,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) { @@ -3458,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}; @@ -3513,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. @@ -3562,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 @@ -3586,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(), @@ -3623,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(); @@ -3663,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, @@ -3683,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}}, @@ -3739,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()); @@ -3758,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>{ @@ -3776,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}, @@ -3816,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()); @@ -3829,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>{ @@ -3846,6 +3849,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); @@ -3879,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_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(); } 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; 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/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( 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); 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. diff --git a/third_party/xla/xla/debug_options_flags.cc b/third_party/xla/xla/debug_options_flags.cc index 4ff61a22dd4f24..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); @@ -3331,7 +3332,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 " @@ -3732,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/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), 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)); 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/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/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) { 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() 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 21b539b21b6e10..b256e1bb8bd482 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()) { @@ -974,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(); @@ -1539,7 +1545,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); + } } } @@ -1660,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(); @@ -2174,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(); @@ -2262,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) { 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_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_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)); 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/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(), 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; diff --git a/third_party/xla/xla/xla.proto b/third_party/xla/xla/xla.proto index 4a56ebe77e5e97..be9d911eff5dd8 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. @@ -1142,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; @@ -1883,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.