From c057b93aa33b4c2ae470246f28da2d70d7c4b1a7 Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Wed, 29 Jul 2026 12:21:00 +0100 Subject: [PATCH 01/15] Add dynamic size expression tests and fix propagation gaps (#27) --- .../tf2xla/kernels/batchtospace_op.cc | 32 +- .../compiler/tf2xla/kernels/bincount_op.cc | 13 +- .../tf2xla/kernels/depthtospace_op.cc | 41 +- tensorflow/compiler/tf2xla/kernels/diag_op.cc | 43 +- .../tf2xla/kernels/dynamic_partition_op.cc | 4 +- .../tf2xla/kernels/dynamic_stitch_op.cc | 10 +- .../tf2xla/kernels/reverse_sequence_op.cc | 27 +- tensorflow/compiler/tf2xla/kernels/roll_op.cc | 4 +- .../compiler/tf2xla/kernels/sequence_ops.cc | 2 +- .../tf2xla/kernels/spacetobatch_op.cc | 23 +- .../tf2xla/kernels/spacetodepth_op.cc | 41 +- .../tf2xla/kernels/strided_slice_op.cc | 21 + .../tf2xla/kernels/tensor_list_utils.cc | 14 +- .../compiler/tf2xla/kernels/unique_op.cc | 12 +- .../compiler/tf2xla/kernels/unpack_op.cc | 9 +- .../compiler/tf2xla/kernels/where_op.cc | 28 +- tensorflow/compiler/tf2xla/shape_util.cc | 3 +- tensorflow/compiler/tf2xla/xla_compiler.cc | 19 +- .../compiler/tf2xla/xla_compiler_test.cc | 2890 ++++++++++++++++- tensorflow/core/framework/tensor_shape.cc | 16 +- .../core/framework/tensor_shape_expr.cc | 12 + tensorflow/core/framework/tensor_shape_expr.h | 5 + .../xla/xla/service/shape_inference.cc | 62 +- .../xla/xla/service/shape_inference_test.cc | 820 +++++ third_party/xla/xla/shape_expr.cc | 21 + 25 files changed, 3945 insertions(+), 227 deletions(-) diff --git a/tensorflow/compiler/tf2xla/kernels/batchtospace_op.cc b/tensorflow/compiler/tf2xla/kernels/batchtospace_op.cc index 7a42150f3a9c19..6deeaa7ba11037 100644 --- a/tensorflow/compiler/tf2xla/kernels/batchtospace_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/batchtospace_op.cc @@ -36,6 +36,8 @@ void BatchToSpace(XlaOpKernelContext* ctx, const xla::XlaOp input, const int input_rank = input_tensor_shape.dims(); const absl::InlinedVector input_shape = input_tensor_shape.dim_sizes(); + const std::vector input_exprs = + input_tensor_shape.get_filled_expressions(); const int block_rank = block_shape.size(); OP_REQUIRES( @@ -76,11 +78,21 @@ void BatchToSpace(XlaOpKernelContext* ctx, const xla::XlaOp input, ") is not divisible by product of block sizes (", block_num_elems, ")")); std::vector reshaped_shape(input_rank + block_rank); + std::vector reshaped_exprs(input_rank + block_rank); std::copy(block_shape.begin(), block_shape.end(), reshaped_shape.begin()); + std::fill(reshaped_exprs.begin(), reshaped_exprs.begin() + block_rank, + xla::DExpr::Const(0)); + for (int i = 0; i < block_rank; ++i) { + reshaped_exprs[i] = xla::DExpr::Const(block_shape[i]); + } reshaped_shape[block_rank] = batch_size / block_num_elems; + reshaped_exprs[block_rank] = + (input_exprs[0] / xla::DExpr::Const(block_num_elems)).simplify(); std::copy(input_shape.begin() + 1, input_shape.end(), reshaped_shape.begin() + block_rank + 1); - xla::XlaOp reshaped = xla::Reshape(input, reshaped_shape); + std::copy(input_exprs.begin() + 1, input_exprs.end(), + reshaped_exprs.begin() + block_rank + 1); + xla::XlaOp reshaped = xla::Reshape(input, reshaped_shape, reshaped_exprs); // 2. Permute dimensions of `reshaped` to produce `permuted` of shape // [batch / prod(block_shape), @@ -111,15 +123,22 @@ void BatchToSpace(XlaOpKernelContext* ctx, const xla::XlaOp input, // ..., // input_shape[N-1]] std::vector reshaped_permuted_shape(input_rank); + std::vector reshaped_permuted_exprs(input_rank); reshaped_permuted_shape[0] = batch_size / block_num_elems; + reshaped_permuted_exprs[0] = + (input_exprs[0] / xla::DExpr::Const(block_num_elems)).simplify(); for (int i = 0; i < block_rank; ++i) { reshaped_permuted_shape[1 + i] = block_shape[i] * input_shape[1 + i]; + reshaped_permuted_exprs[1 + i] = + (xla::DExpr::Const(block_shape[i]) * input_exprs[1 + i]).simplify(); } std::copy(remainder_shape.begin(), remainder_shape.end(), reshaped_permuted_shape.begin() + 1 + block_rank); + std::copy(input_exprs.begin() + 1 + block_rank, input_exprs.end(), + reshaped_permuted_exprs.begin() + 1 + block_rank); xla::XlaOp reshaped_permuted = - xla::Reshape(permuted, reshaped_permuted_shape); + xla::Reshape(permuted, reshaped_permuted_shape, reshaped_permuted_exprs); // 4. Crop the start and end of dimensions `[1, ..., M]` of // `reshaped_permuted` according to `crops` to produce the output of shape: @@ -133,6 +152,9 @@ void BatchToSpace(XlaOpKernelContext* ctx, const xla::XlaOp input, std::vector start_indices(input_rank, 0); std::vector end_indices = reshaped_permuted_shape; std::vector strides(input_rank, 1); + std::vector start_exprs(input_rank, xla::DExpr::Const(0)); + std::vector end_exprs(reshaped_permuted_exprs.begin(), + reshaped_permuted_exprs.end()); for (int i = 0; i < block_rank; ++i) { int64_t crop_start = crops.Get({i, 0}); int64_t crop_end = crops.Get({i, 1}); @@ -140,14 +162,16 @@ void BatchToSpace(XlaOpKernelContext* ctx, const xla::XlaOp input, errors::InvalidArgument("Crops must be non-negative")); start_indices[1 + i] = crop_start; end_indices[1 + i] -= crop_end; + start_exprs[1 + i] = xla::DExpr::Const(crop_start); + end_exprs[1 + i] = (reshaped_permuted_exprs[1 + i] - crop_end).simplify(); OP_REQUIRES( ctx, start_indices[1 + i] <= end_indices[1 + i], errors::InvalidArgument( "Cropped size must be non-negative: start: ", crop_start, " end: ", crop_end, " size ", reshaped_permuted_shape[1 + i])); } - xla::XlaOp output = - xla::Slice(reshaped_permuted, start_indices, end_indices, strides); + xla::XlaOp output = xla::Slice(reshaped_permuted, start_indices, end_indices, + start_exprs, end_exprs, strides); ctx->SetOutput(0, output); } diff --git a/tensorflow/compiler/tf2xla/kernels/bincount_op.cc b/tensorflow/compiler/tf2xla/kernels/bincount_op.cc index df3347bd5d533f..5e277772f40e4e 100644 --- a/tensorflow/compiler/tf2xla/kernels/bincount_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/bincount_op.cc @@ -110,11 +110,15 @@ class DenseBincountOp : public XlaOpKernel { scatter_dnums.add_scatter_dims_to_operand_dims(0); if (rank == 2) { - output_shape = xla::ShapeUtil::MakeShape(dtype, {size, output_size}); + output_shape = xla::ShapeUtil::MakeShape( + dtype, {size, output_size}, + std::vector{input_shape.expressions(0), + xla::DExpr::Const(output_size)}); scatter_dnums.add_inserted_window_dims(1); scatter_dnums.add_scatter_dims_to_operand_dims(1); - auto i_shape = - xla::ShapeUtil::MakeShape(input_xla_type, {input_shape.dimensions()}); + auto i_shape = xla::ShapeUtil::MakeShape(input_xla_type, + input_shape.dimensions(), + input_shape.expressions()); auto i = xla::Iota(ctx->builder(), i_shape, 0); xla::DExpr flattened_expr = input_shape.expressions(0) * input_shape.expressions(1); @@ -131,7 +135,8 @@ class DenseBincountOp : public XlaOpKernel { updates = xla::Broadcast( one, {input_shape.dimensions(0) * input_shape.dimensions(1)}); output = xla::Broadcast( - zero, {output_shape.dimensions(0), output_shape.dimensions(1)}); + zero, {output_shape.dimensions(0), output_shape.dimensions(1)}, + {output_shape.expressions(0), output_shape.expressions(1)}); if (has_weights && !binary_output_) { weights = xla::Reshape( weights, {input_shape.dimensions(0) * input_shape.dimensions(1)}, diff --git a/tensorflow/compiler/tf2xla/kernels/depthtospace_op.cc b/tensorflow/compiler/tf2xla/kernels/depthtospace_op.cc index e8e2babffd529c..3a651d1adae655 100644 --- a/tensorflow/compiler/tf2xla/kernels/depthtospace_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/depthtospace_op.cc @@ -64,6 +64,7 @@ class DepthToSpaceOp : public XlaOpKernel { OP_REQUIRES_OK(ctx, input_xla_shape.status()); absl::Span input_shape = input_xla_shape.value().dimensions(); + const xla::Shape& input_shape_with_exprs = input_xla_shape.value(); int input_rank = input_shape.size(); static const int kRequiredDims = 4; @@ -77,20 +78,31 @@ class DepthToSpaceOp : public XlaOpKernel { std::vector reshaped_shape; std::vector transpose_order; std::vector output_shape; + std::vector reshaped_exprs; + std::vector output_exprs; reshaped_shape.reserve(input_rank); transpose_order.reserve(input_rank); output_shape.reserve(input_rank); + reshaped_exprs.reserve(input_rank + num_spatial_dims); + output_exprs.reserve(input_rank); if (data_format == FORMAT_NHWC) { reshaped_shape.push_back(input_shape[0]); + reshaped_exprs.push_back(input_shape_with_exprs.expressions(0)); for (int i = 0; i < num_spatial_dims; ++i) { reshaped_shape.push_back(input_shape[1 + i]); + reshaped_exprs.push_back(input_shape_with_exprs.expressions(1 + i)); } int64_t block_elems = 1; for (int i = 0; i < num_spatial_dims; ++i) { reshaped_shape.push_back(block_size_); + reshaped_exprs.push_back(xla::DExpr::Const(block_size_)); block_elems *= block_size_; } reshaped_shape.push_back(input_shape[feature_dim] / block_elems); + reshaped_exprs.push_back( + (input_shape_with_exprs.expressions(feature_dim) / + xla::DExpr::Const(block_elems)) + .simplify()); transpose_order.push_back(0); for (int i = 0; i < num_spatial_dims; ++i) { @@ -100,21 +112,37 @@ class DepthToSpaceOp : public XlaOpKernel { transpose_order.push_back(feature_dim + num_spatial_dims); output_shape.push_back(input_shape[0]); + output_exprs.push_back(input_shape_with_exprs.expressions(0)); for (int i = 0; i < num_spatial_dims; ++i) { output_shape.push_back(input_shape[1 + i] * block_size_); + output_exprs.push_back( + (input_shape_with_exprs.expressions(1 + i) * + xla::DExpr::Const(block_size_)) + .simplify()); } output_shape.push_back(input_shape[feature_dim] / block_elems); + output_exprs.push_back( + (input_shape_with_exprs.expressions(feature_dim) / + xla::DExpr::Const(block_elems)) + .simplify()); } else { // NCHW format. reshaped_shape.push_back(input_shape[0]); + reshaped_exprs.push_back(input_shape_with_exprs.expressions(0)); int64_t block_elems = 1; for (int i = 0; i < num_spatial_dims; ++i) { reshaped_shape.push_back(block_size_); + reshaped_exprs.push_back(xla::DExpr::Const(block_size_)); block_elems *= block_size_; } reshaped_shape.push_back(input_shape[feature_dim] / block_elems); + reshaped_exprs.push_back( + (input_shape_with_exprs.expressions(feature_dim) / + xla::DExpr::Const(block_elems)) + .simplify()); for (int i = 0; i < num_spatial_dims; ++i) { reshaped_shape.push_back(input_shape[2 + i]); + reshaped_exprs.push_back(input_shape_with_exprs.expressions(2 + i)); } transpose_order.push_back(0); @@ -125,9 +153,18 @@ class DepthToSpaceOp : public XlaOpKernel { } output_shape.push_back(input_shape[0]); + output_exprs.push_back(input_shape_with_exprs.expressions(0)); output_shape.push_back(input_shape[feature_dim] / block_elems); + output_exprs.push_back( + (input_shape_with_exprs.expressions(feature_dim) / + xla::DExpr::Const(block_elems)) + .simplify()); for (int i = 0; i < num_spatial_dims; ++i) { output_shape.push_back(input_shape[2 + i] * block_size_); + output_exprs.push_back( + (input_shape_with_exprs.expressions(2 + i) * + xla::DExpr::Const(block_size_)) + .simplify()); } } @@ -148,7 +185,7 @@ class DepthToSpaceOp : public XlaOpKernel { ") is not divisible by square of the block size (", block_size_, ")")); - xla::XlaOp reshaped = xla::Reshape(input, reshaped_shape); + xla::XlaOp reshaped = xla::Reshape(input, reshaped_shape, reshaped_exprs); // 2. Permute dimensions of `reshaped` to produce // `permuted_reshaped` of shape: @@ -169,7 +206,7 @@ class DepthToSpaceOp : public XlaOpKernel { // input_shape[2] * block_size_, // depth / (block_size_ * block_size_)] // - xla::XlaOp output = xla::Reshape(permuted_reshaped, output_shape); + xla::XlaOp output = xla::Reshape(permuted_reshaped, output_shape, output_exprs); // If this used to be a vectorized format turn it back now. if (data_format != data_format_) { diff --git a/tensorflow/compiler/tf2xla/kernels/diag_op.cc b/tensorflow/compiler/tf2xla/kernels/diag_op.cc index d8740aa137fbeb..9644d52fa489a6 100644 --- a/tensorflow/compiler/tf2xla/kernels/diag_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/diag_op.cc @@ -29,6 +29,7 @@ limitations under the License. #include "xla/hlo/builder/lib/matrix.h" #include "xla/hlo/builder/lib/pooling.h" #include "xla/hlo/builder/xla_builder.h" +#include "xla/shape_util.h" #include "xla/util.h" #include "xla/xla_data.pb.h" #include "tensorflow/core/framework/op_kernel.h" @@ -38,7 +39,9 @@ namespace { // Create a diagonal / batch diagonal matrix with 'input' on the diagonal. xla::XlaOp CreateDiagonal(xla::XlaOp input, int64_t last_dim_size, - absl::Span other_dims) { + const xla::DExpr& last_dim_expr, + absl::Span other_dims, + absl::Span other_dim_exprs) { xla::XlaBuilder* builder = input.builder(); // Create two matrices that have the following forms, and compare them: // @@ -49,14 +52,23 @@ xla::XlaOp CreateDiagonal(xla::XlaOp input, int64_t last_dim_size, // // This produces a predicate matrix of the right size, with "true" on the // diagonal. - xla::XlaOp iota = xla::Iota(builder, xla::S32, last_dim_size); - xla::XlaOp iota_broadcast = xla::Broadcast(iota, {last_dim_size}); + xla::XlaOp iota = xla::Iota( + builder, + xla::ShapeUtil::MakeShape(xla::S32, std::vector{last_dim_size}, + std::vector{last_dim_expr}), + /*iota_dimension=*/0); + xla::XlaOp iota_broadcast = xla::Broadcast( + iota, {last_dim_size}, {last_dim_expr, last_dim_expr}); xla::XlaOp mask = xla::Eq(iota_broadcast, iota, {0}); // If this is a batched diagonal, broadcast the mask across the other // dimensions. if (!other_dims.empty()) { - mask = xla::Broadcast(mask, other_dims); + std::vector mask_exprs(other_dim_exprs.begin(), + other_dim_exprs.end()); + mask_exprs.push_back(last_dim_expr); + mask_exprs.push_back(last_dim_expr); + mask = xla::Broadcast(mask, other_dims, mask_exprs); } // Broadcast the input, and then use the mask computed above to select the @@ -69,13 +81,17 @@ xla::XlaOp CreateDiagonal(xla::XlaOp input, int64_t last_dim_size, std::vector out_dim_sizes(other_dims.begin(), other_dims.end()); out_dim_sizes.push_back(last_dim_size); out_dim_sizes.push_back(last_dim_size); + std::vector out_dim_exprs(other_dim_exprs.begin(), + other_dim_exprs.end()); + out_dim_exprs.push_back(last_dim_expr); + out_dim_exprs.push_back(last_dim_expr); // Broadcast into the second to last dimension. std::vector broadcast_dimensions(other_dims.size() + 1); absl::c_iota(broadcast_dimensions, 0); ++broadcast_dimensions.back(); - xla::XlaOp input_broadcast = - xla::BroadcastInDim(input, out_dim_sizes, broadcast_dimensions); + xla::XlaOp input_broadcast = xla::BroadcastInDim( + input, out_dim_sizes, broadcast_dimensions, out_dim_exprs); return xla::Select(mask, input_broadcast, xla::ZerosLike(input_broadcast)); } @@ -102,17 +118,26 @@ class DiagOp : public XlaOpKernel { // [0, 0, 0, 4]] // Flattens the input to 1D. + xla::DExpr flattened_expr = xla::DExpr::Const(1); + std::vector input_exprs = input_shape.get_filled_expressions(); + for (const xla::DExpr& expr : input_exprs) { + flattened_expr = (flattened_expr * expr).simplify(); + } int64_t size = input_shape.num_elements(); - input = xla::Reshape(input, {size}, {}); + input = xla::Reshape(input, {size}, {flattened_expr}); // Create an R2 with the R1 diagonal. - xla::XlaOp diag = CreateDiagonal(input, size, /*other_dims=*/{}); + xla::XlaOp diag = + CreateDiagonal(input, size, flattened_expr, /*other_dims=*/{}, + /*other_dim_exprs=*/{}); // Reshapes to the final shape. std::vector new_dims(dims.size() * 2); std::copy(dims.begin(), dims.end(), new_dims.begin()); std::copy(dims.begin(), dims.end(), new_dims.begin() + dims.size()); - diag = xla::Reshape(diag, new_dims); + std::vector new_exprs(input_exprs.begin(), input_exprs.end()); + new_exprs.insert(new_exprs.end(), input_exprs.begin(), input_exprs.end()); + diag = xla::Reshape(diag, new_dims, new_exprs); ctx->SetOutput(0, diag); } diff --git a/tensorflow/compiler/tf2xla/kernels/dynamic_partition_op.cc b/tensorflow/compiler/tf2xla/kernels/dynamic_partition_op.cc index 85f19f9541033b..2b9f70ffecc681 100644 --- a/tensorflow/compiler/tf2xla/kernels/dynamic_partition_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/dynamic_partition_op.cc @@ -213,11 +213,11 @@ class DynamicPartitionOp : public XlaOpKernel { {CollapseExpressions(flattened_partition_exprs)}); xla::Shape data_1d_shape = xla::ShapeUtil::MakeShape( data_shape.element_type(), {input_count}, - {xla::DExpr::Const(input_count)}); + std::vector{xla::DExpr::Const(input_count)}); xla::Shape partitions_1d_shape = xla::ShapeUtil::MakeShape( partition_shape.element_type(), {input_count}, - {xla::DExpr::Const(input_count)}); + std::vector{xla::DExpr::Const(input_count)}); std::vector output, partition_length; std::tie(output, partition_length) = DynamicPartition1D( diff --git a/tensorflow/compiler/tf2xla/kernels/dynamic_stitch_op.cc b/tensorflow/compiler/tf2xla/kernels/dynamic_stitch_op.cc index 305b527cc76632..fe45493d452d46 100644 --- a/tensorflow/compiler/tf2xla/kernels/dynamic_stitch_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/dynamic_stitch_op.cc @@ -132,13 +132,18 @@ class DynamicStitchOp : public XlaOpKernel { int64_t result_rank = 1 + data0_shape.dims() - indices0_shape.dims(); if (number_of_indices == 0) { std::vector result_shape(result_rank); + std::vector result_expressions(result_rank, + xla::DExpr::Const(0)); for (int d = indices0_shape.dims(); d < data0_shape.dims(); d++) { result_shape[d - indices0_shape.dims() + 1] = data0_shape.dim_size(d); + result_expressions[d - indices0_shape.dims() + 1] = + data0_shape.get_filled_expression(d); } xla::PrimitiveType element_type = ctx->input_xla_type(ctx->num_inputs() - 1); xla::Literal empty_literal = xla::Literal::CreateFromShape( - xla::ShapeUtil::MakeShape(element_type, result_shape)); + xla::ShapeUtil::MakeShape(element_type, result_shape, + result_expressions)); ctx->SetOutput(0, xla::ConstantLiteral(ctx->builder(), empty_literal)); return; } @@ -186,7 +191,8 @@ class DynamicStitchOp : public XlaOpKernel { if (new_shape == data_shapes[input_num]) { input[input_num] = handle; } else { - input[input_num] = xla::Reshape(handle, new_shape.dim_sizes()); + input[input_num] = xla::Reshape(handle, new_shape.dim_sizes(), + new_shape.get_filled_expressions()); } } diff --git a/tensorflow/compiler/tf2xla/kernels/reverse_sequence_op.cc b/tensorflow/compiler/tf2xla/kernels/reverse_sequence_op.cc index cb6f8cebf0a8d9..9f3bbf333eb4eb 100644 --- a/tensorflow/compiler/tf2xla/kernels/reverse_sequence_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/reverse_sequence_op.cc @@ -71,6 +71,9 @@ class ReverseSequenceOp : public XlaOpKernel { xla::XlaBuilder* builder = context->builder(); const auto input = context->Input(0); const auto seq_lens = context->Input(1); + auto input_xla_shape_or = context->InputXlaShape(0); + OP_REQUIRES(context, input_xla_shape_or.ok(), input_xla_shape_or.status()); + const xla::Shape& input_xla_shape = input_xla_shape_or.value(); const int64_t batch_size = input_shape.dim_size(batch_dim_); if (batch_size == 0) { @@ -86,17 +89,21 @@ class ReverseSequenceOp : public XlaOpKernel { xla::XlaOp back = xla::Sub(seq_lens, xla::ScalarLike(seq_lens, 1)); xla::XlaOp batch_idx = xla::Iota( builder, - xla::ShapeUtil::MakeShape(seq_lens_type, {batch_size, max_seq_len, 1}, - {input_shape.get_filled_expression(batch_dim_), - input_shape.get_filled_expression(seq_dim_), - xla::DExpr::Const(1)}), + xla::ShapeUtil::MakeShape( + seq_lens_type, {batch_size, max_seq_len, 1}, + std::vector{ + input_shape.get_filled_expression(batch_dim_), + input_shape.get_filled_expression(seq_dim_), + xla::DExpr::Const(1)}), /*iota_dimension=*/0); xla::XlaOp forward_idx = xla::Iota( builder, - xla::ShapeUtil::MakeShape(seq_lens_type, {batch_size, max_seq_len, 1}, - {input_shape.get_filled_expression(batch_dim_), - input_shape.get_filled_expression(seq_dim_), - xla::DExpr::Const(1)}), + xla::ShapeUtil::MakeShape( + seq_lens_type, {batch_size, max_seq_len, 1}, + std::vector{ + input_shape.get_filled_expression(batch_dim_), + input_shape.get_filled_expression(seq_dim_), + xla::DExpr::Const(1)}), /*iota_dimension=*/1); xla::XlaOp reverse_idx = xla::Sub(back, forward_idx, {0}); reverse_idx = xla::Select(xla::Lt(reverse_idx, xla::ZerosLike(reverse_idx)), @@ -135,8 +142,8 @@ class ReverseSequenceOp : public XlaOpKernel { slice_sizes[batch_dim_] = 1; slice_sizes[seq_dim_] = 1; - context->SetOutput(0, - xla::Gather(input, start_indices, dnums, slice_sizes)); + xla::XlaOp gathered = xla::Gather(input, start_indices, dnums, slice_sizes); + context->SetOutput(0, xla::Reshape(input_xla_shape, gathered)); } private: diff --git a/tensorflow/compiler/tf2xla/kernels/roll_op.cc b/tensorflow/compiler/tf2xla/kernels/roll_op.cc index 0fcc6bec56095b..49b5cd3f01b8b8 100644 --- a/tensorflow/compiler/tf2xla/kernels/roll_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/roll_op.cc @@ -94,8 +94,8 @@ class RollOp : public XlaOpKernel { std::vector start_indices( input_shape.dims(), xla::Zero(ctx->builder(), shift_type)); start_indices[cur_axis] = axis_size - offset; - output = - xla::DynamicSlice(concat, start_indices, input_shape.dim_sizes()); + output = xla::DynamicSlice(concat, start_indices, input_shape.dim_sizes(), + input_shape.get_filled_expressions()); } ctx->SetOutput(0, output); } diff --git a/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc b/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc index 8f8a34899fc46a..1ff48ac5d64a66 100644 --- a/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc +++ b/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc @@ -143,7 +143,7 @@ absl::StatusOr CreateRangeTensor( ? xla::Iota(builder, xla::ShapeUtil::MakeShape( xla::primitive_util::NativeToPrimitiveType(), - {size}, {size_expr}), + {size}, std::vector{size_expr}), /*iota_dimension=*/0) : xla::Iota(builder, xla::primitive_util::NativeToPrimitiveType(), size); diff --git a/tensorflow/compiler/tf2xla/kernels/spacetobatch_op.cc b/tensorflow/compiler/tf2xla/kernels/spacetobatch_op.cc index d4a93e0556143d..e10c009a5b7d16 100644 --- a/tensorflow/compiler/tf2xla/kernels/spacetobatch_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/spacetobatch_op.cc @@ -44,6 +44,8 @@ void SpaceToBatch(XlaOpKernelContext* ctx, const xla::XlaOp input, const int input_rank = input_tensor_shape.dims(); const absl::InlinedVector input_shape = input_tensor_shape.dim_sizes(); + const std::vector input_exprs = + input_tensor_shape.get_filled_expressions(); const int block_rank = block_shape.size(); OP_REQUIRES( @@ -68,6 +70,7 @@ void SpaceToBatch(XlaOpKernelContext* ctx, const xla::XlaOp input, // input according to `paddings` to produce `padded` of shape `padded_shape`. xla::PaddingConfig padding_config; std::vector padded_shape(input_shape.begin(), input_shape.end()); + std::vector padded_exprs(input_exprs.begin(), input_exprs.end()); int64_t block_num_elems = 1LL; padding_config.add_dimensions(); // Don't pad the batch dimension. for (int i = 0; i < block_rank; ++i) { @@ -83,6 +86,7 @@ void SpaceToBatch(XlaOpKernelContext* ctx, const xla::XlaOp input, dim->set_edge_padding_low(pad_start); dim->set_edge_padding_high(pad_end); padded_shape[1 + i] += pad_start + pad_end; + padded_exprs[1 + i] = (padded_exprs[1 + i] + pad_start + pad_end).simplify(); block_num_elems = MultiplyWithoutOverflow(block_num_elems, block_shape[i]); } // Don't pad the remainder dimensions. @@ -116,7 +120,9 @@ void SpaceToBatch(XlaOpKernelContext* ctx, const xla::XlaOp input, // block_shape[M-1]] + // remaining_shape std::vector reshaped_padded_shape(input_rank + block_rank); + std::vector reshaped_padded_exprs(input_rank + block_rank); reshaped_padded_shape[0] = batch_size; + reshaped_padded_exprs[0] = padded_exprs[0]; for (int i = 0; i < block_rank; ++i) { OP_REQUIRES(ctx, padded_shape[1 + i] % block_shape[i] == 0, errors::InvalidArgument("padded_shape[", 1 + i, @@ -126,11 +132,17 @@ void SpaceToBatch(XlaOpKernelContext* ctx, const xla::XlaOp input, reshaped_padded_shape[1 + i * 2] = padded_shape[1 + i] / block_shape[i]; reshaped_padded_shape[1 + i * 2 + 1] = block_shape[i]; + reshaped_padded_exprs[1 + i * 2] = + (padded_exprs[1 + i] / block_shape[i]).simplify(); + reshaped_padded_exprs[1 + i * 2 + 1] = xla::DExpr::Const(block_shape[i]); } std::copy(remainder_shape.begin(), remainder_shape.end(), reshaped_padded_shape.begin() + 1 + 2 * block_rank); + std::copy(input_exprs.begin() + 1 + block_rank, input_exprs.end(), + reshaped_padded_exprs.begin() + 1 + 2 * block_rank); - xla::XlaOp reshaped_padded = xla::Reshape(padded, reshaped_padded_shape); + xla::XlaOp reshaped_padded = + xla::Reshape(padded, reshaped_padded_shape, reshaped_padded_exprs); // 3. Permute dimensions of `reshaped_padded` to produce // `permuted_reshaped_padded` of shape: @@ -163,14 +175,21 @@ void SpaceToBatch(XlaOpKernelContext* ctx, const xla::XlaOp input, // Determine the length of the prefix of block dims that can be combined // into the batch dimension due to having no padding and block_shape=1. std::vector output_shape(input_rank); + std::vector output_exprs(input_rank); output_shape[0] = output_dim; + output_exprs[0] = (input_exprs[0] * xla::DExpr::Const(block_num_elems)).simplify(); for (int i = 0; i < block_rank; ++i) { output_shape[1 + i] = padded_shape[1 + i] / block_shape[i]; + output_exprs[1 + i] = + (padded_exprs[1 + i] / block_shape[i]).simplify(); } std::copy(remainder_shape.begin(), remainder_shape.end(), output_shape.begin() + 1 + block_rank); + std::copy(input_exprs.begin() + 1 + block_rank, input_exprs.end(), + output_exprs.begin() + 1 + block_rank); - xla::XlaOp output = xla::Reshape(permuted_reshaped_padded, output_shape); + xla::XlaOp output = + xla::Reshape(permuted_reshaped_padded, output_shape, output_exprs); ctx->SetOutput(0, output); } diff --git a/tensorflow/compiler/tf2xla/kernels/spacetodepth_op.cc b/tensorflow/compiler/tf2xla/kernels/spacetodepth_op.cc index ac33e0877200dc..d09fa5f4daacbb 100644 --- a/tensorflow/compiler/tf2xla/kernels/spacetodepth_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/spacetodepth_op.cc @@ -67,6 +67,7 @@ class SpaceToDepthOp : public XlaOpKernel { OP_REQUIRES_OK(ctx, input_xla_shape.status()); absl::Span input_shape = input_xla_shape.value().dimensions(); + const xla::Shape& input_shape_with_exprs = input_xla_shape.value(); int input_rank = input_shape.size(); static const int kRequiredDims = 4; @@ -80,9 +81,13 @@ class SpaceToDepthOp : public XlaOpKernel { std::vector reshaped_shape; std::vector transpose_order; std::vector output_shape; + std::vector reshaped_exprs; + std::vector output_exprs; reshaped_shape.reserve(input_rank); transpose_order.reserve(input_rank); output_shape.reserve(input_rank); + reshaped_exprs.reserve(input_rank + num_spatial_dims); + output_exprs.reserve(input_rank); if (data_format == FORMAT_NHWC) { int64_t block_elems = 1; for (int i = 0; i < num_spatial_dims; ++i) { @@ -94,11 +99,18 @@ class SpaceToDepthOp : public XlaOpKernel { } reshaped_shape.push_back(input_shape[0]); + reshaped_exprs.push_back(input_shape_with_exprs.expressions(0)); for (int i = 0; i < num_spatial_dims; ++i) { reshaped_shape.push_back(input_shape[1 + i] / block_size_); + reshaped_exprs.push_back( + (input_shape_with_exprs.expressions(1 + i) / + xla::DExpr::Const(block_size_)) + .simplify()); reshaped_shape.push_back(block_size_); + reshaped_exprs.push_back(xla::DExpr::Const(block_size_)); } reshaped_shape.push_back(input_shape[feature_dim]); + reshaped_exprs.push_back(input_shape_with_exprs.expressions(feature_dim)); transpose_order.push_back(0); for (int i = 0; i < num_spatial_dims; ++i) { @@ -110,10 +122,19 @@ class SpaceToDepthOp : public XlaOpKernel { transpose_order.push_back(feature_dim + num_spatial_dims); output_shape.push_back(input_shape[0]); + output_exprs.push_back(input_shape_with_exprs.expressions(0)); for (int i = 0; i < num_spatial_dims; ++i) { output_shape.push_back(input_shape[1 + i] / block_size_); + output_exprs.push_back( + (input_shape_with_exprs.expressions(1 + i) / + xla::DExpr::Const(block_size_)) + .simplify()); } output_shape.push_back(input_shape[feature_dim] * block_elems); + output_exprs.push_back( + (input_shape_with_exprs.expressions(feature_dim) * + xla::DExpr::Const(block_elems)) + .simplify()); } else { // FORMAT_NCHW int64_t block_elems = 1; @@ -126,10 +147,17 @@ class SpaceToDepthOp : public XlaOpKernel { } reshaped_shape.push_back(input_shape[0]); + reshaped_exprs.push_back(input_shape_with_exprs.expressions(0)); reshaped_shape.push_back(input_shape[feature_dim]); + reshaped_exprs.push_back(input_shape_with_exprs.expressions(feature_dim)); for (int i = 0; i < num_spatial_dims; ++i) { reshaped_shape.push_back(input_shape[2 + i] / block_size_); + reshaped_exprs.push_back( + (input_shape_with_exprs.expressions(2 + i) / + xla::DExpr::Const(block_size_)) + .simplify()); reshaped_shape.push_back(block_size_); + reshaped_exprs.push_back(xla::DExpr::Const(block_size_)); } transpose_order.push_back(0); @@ -142,9 +170,18 @@ class SpaceToDepthOp : public XlaOpKernel { } output_shape.push_back(input_shape[0]); + output_exprs.push_back(input_shape_with_exprs.expressions(0)); output_shape.push_back(input_shape[feature_dim] * block_elems); + output_exprs.push_back( + (input_shape_with_exprs.expressions(feature_dim) * + xla::DExpr::Const(block_elems)) + .simplify()); for (int i = 0; i < num_spatial_dims; ++i) { output_shape.push_back(input_shape[2 + i] / block_size_); + output_exprs.push_back( + (input_shape_with_exprs.expressions(2 + i) / + xla::DExpr::Const(block_size_)) + .simplify()); } } @@ -156,7 +193,7 @@ class SpaceToDepthOp : public XlaOpKernel { // input_shape[1] / block_size_, block_size_, // input_shape[2] / block_size_, block_size_, // depth] - xla::XlaOp reshaped = xla::Reshape(input, reshaped_shape); + xla::XlaOp reshaped = xla::Reshape(input, reshaped_shape, reshaped_exprs); // 2. Permute dimensions of `reshaped` to produce // `permuted_reshaped` of shape: @@ -176,7 +213,7 @@ class SpaceToDepthOp : public XlaOpKernel { // input_shape[2] / block_size_, // block_size_ * block_size_ * depth] // - xla::XlaOp output = xla::Reshape(permuted_reshaped, output_shape); + xla::XlaOp output = xla::Reshape(permuted_reshaped, output_shape, output_exprs); // If this used to be a vectorized format turn it back now. if (data_format != data_format_) { diff --git a/tensorflow/compiler/tf2xla/kernels/strided_slice_op.cc b/tensorflow/compiler/tf2xla/kernels/strided_slice_op.cc index 66c369b70f7591..9189a83d3f7f69 100644 --- a/tensorflow/compiler/tf2xla/kernels/strided_slice_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/strided_slice_op.cc @@ -373,6 +373,27 @@ class StridedSliceOp : public XlaOpKernel { if (!dimensions_to_reverse.empty()) { slice = xla::Rev(slice, dimensions_to_reverse); } + for (int i = 0; i < partial_processing_shape.dims(); ++i) { + partial_processing_shape.set_expression( + i, ((slice_end_expr[i] - slice_begin_expr[i] + + xla::DExpr::Const(slice_strides[i]) - xla::DExpr::Const(1)) / + xla::DExpr::Const(slice_strides[i])) + .simplify()); + } + for (int i = 0; i < partial_final_shape.dims(); ++i) { + int64_t processing_index = shape_spec.output_to_processing_mapping[i]; + partial_final_shape.set_expression( + i, processing_index == -1 + ? xla::DExpr::Const(partial_final_shape.dim_size(i)) + : partial_processing_shape.get_filled_expression( + processing_index)); + } + OP_REQUIRES( + ctx, partial_final_shape.AsTensorShape(&final_shape), + InvalidArgument("XLA can't deduce compile time constant output " + "shape for strided slice: ", + partial_final_shape.DebugString(), + ", output shape must be a compile-time constant")); slice = enable_dynamic_sizes ? xla::Slice(slice, slice_begin, slice_end, slice_begin_expr, slice_end_expr, slice_strides) diff --git a/tensorflow/compiler/tf2xla/kernels/tensor_list_utils.cc b/tensorflow/compiler/tf2xla/kernels/tensor_list_utils.cc index 9f235e6994e7d9..412c01f24dfc19 100644 --- a/tensorflow/compiler/tf2xla/kernels/tensor_list_utils.cc +++ b/tensorflow/compiler/tf2xla/kernels/tensor_list_utils.cc @@ -241,9 +241,12 @@ absl::Status GetTensorListShapeFromElementTensorListShape( const xla::Shape& shape = xla::ShapeUtil::GetTupleElementShape(element_tensor_list_shape, i); std::vector dimensions = xla::SpanToVector(shape.dimensions()); + std::vector expressions = xla::SpanToVector(shape.expressions()); dimensions.insert(dimensions.begin(), leading_dim); + expressions.insert(expressions.begin(), xla::DExpr::Const(leading_dim)); shapes.push_back( - xla::ShapeUtil::MakeShape(shape.element_type(), dimensions)); + xla::ShapeUtil::MakeShape(shape.element_type(), dimensions, + expressions)); if (leading_dim_is_dynamic) { shapes.back().set_dynamic_dimension(0, true); } @@ -267,9 +270,13 @@ absl::Status GetTensorListShapeFromElementShape(const xla::Shape& element_shape, std::vector shapes; std::vector dimensions = xla::SpanToVector(element_shape.dimensions()); + std::vector expressions = + xla::SpanToVector(element_shape.expressions()); dimensions.insert(dimensions.begin(), leading_dim); + expressions.insert(expressions.begin(), xla::DExpr::Const(leading_dim)); shapes.push_back( - xla::ShapeUtil::MakeShape(element_shape.element_type(), dimensions)); + xla::ShapeUtil::MakeShape(element_shape.element_type(), dimensions, + expressions)); shapes.back().set_dynamic_dimension(0, leading_dim_is_dynamic); shapes.push_back(xla::ShapeUtil::MakeShape(xla::PrimitiveType::S32, std::vector{})); @@ -289,7 +296,8 @@ absl::Status CreateZerosTensorListWithShape( xla::ShapeUtil::GetTupleElementShape(list_shape, i); xla::XlaOp zero = xla::ConstantLiteral(b, xla::LiteralUtil::Zero(shape.element_type())); - xla::XlaOp zeros = xla::Broadcast(zero, shape.dimensions()); + xla::XlaOp zeros = + xla::Broadcast(zero, shape.dimensions(), shape.expressions()); TF_RET_CHECK(dynamic_dims[i].size() == shape.dimensions().size()); for (int64_t dim = 0; dim < shape.dimensions().size(); ++dim) { if (shape.is_dynamic_dimension(dim)) { diff --git a/tensorflow/compiler/tf2xla/kernels/unique_op.cc b/tensorflow/compiler/tf2xla/kernels/unique_op.cc index f19278265b8ccd..6590dd2f55cf78 100644 --- a/tensorflow/compiler/tf2xla/kernels/unique_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/unique_op.cc @@ -178,13 +178,16 @@ class UniqueOpBase : public XlaOpKernel { sort_keys.reserve(product + 1); std::vector sort_types; sort_types.reserve(product + 1); + xla::Shape leading_shape = xla::ShapeUtil::MakeShape( + input_shape.element_type(), {leading_size}, + std::vector{leading_expr}); for (int64_t i = 0; i < product; ++i) { xla::XlaOp slice = xla::SliceInDim(aux, i, i + 1, 1, 1); - sort_keys.push_back(xla::Reshape(slice, {leading_size}, {leading_expr})); + sort_keys.push_back(xla::Reshape(leading_shape, slice)); sort_types.push_back(input_shape.element_type()); } - xla::Shape iota_shape = - xla::ShapeUtil::MakeShape(xla::S32, {leading_size}, {leading_expr}); + xla::Shape iota_shape = xla::ShapeUtil::MakeShape( + xla::S32, {leading_size}, std::vector{leading_expr}); iota_shape.set_expression(0, leading_expr); auto iota = xla::Iota(ctx->builder(), iota_shape, 0); sort_keys.push_back(iota); @@ -248,8 +251,7 @@ class UniqueOpBase : public XlaOpKernel { /*is_stable=*/true); auto mask_permute = xla::GetTupleElement(mask_sort, 1); permuted = xla::Gather(aux, mask_permute, gather_dim_numbers, {1, product}); - auto result_data = - xla::Reshape(permuted, aux_shape.dimensions(), aux_shape.expressions()); + auto result_data = xla::Reshape(aux_shape, permuted); result_data = MoveAxis(result_data, 0, axis, aux_shape); result_data = xla::SetDimensionSize(result_data, dynamic_size, axis); ctx->SetOutput(0, result_data); diff --git a/tensorflow/compiler/tf2xla/kernels/unpack_op.cc b/tensorflow/compiler/tf2xla/kernels/unpack_op.cc index ee68c3f3aabf5f..55899c7f7b7d95 100644 --- a/tensorflow/compiler/tf2xla/kernels/unpack_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/unpack_op.cc @@ -60,14 +60,21 @@ class UnpackOp : public XlaOpKernel { std::vector start_indices(input_shape.dims(), 0); std::vector limit_indices(input_shape.dims()); std::vector strides(input_shape.dims(), 1); + std::vector start_exprs(input_shape.dims(), xla::DExpr::Const(0)); + std::vector limit_exprs; + limit_exprs.reserve(input_shape.dims()); for (int i = 0; i < input_shape.dims(); ++i) { limit_indices[i] = input_shape.dim_size(i); + limit_exprs.push_back(input_shape.get_filled_expression(i)); } for (int i = 0; i < num; ++i) { start_indices[axis] = i; limit_indices[axis] = i + 1; - auto slice = xla::Slice(input, start_indices, limit_indices, strides); + start_exprs[axis] = xla::DExpr::Const(i); + limit_exprs[axis] = xla::DExpr::Const(i + 1); + auto slice = xla::Slice(input, start_indices, limit_indices, start_exprs, + limit_exprs, strides); // Reshape to drop the 'axis' dimension. auto result = xla::Reshape(slice, output_shape.dim_sizes(), output_shape.get_filled_expressions()); diff --git a/tensorflow/compiler/tf2xla/kernels/where_op.cc b/tensorflow/compiler/tf2xla/kernels/where_op.cc index a83ba478bbb6d7..4920a106808605 100644 --- a/tensorflow/compiler/tf2xla/kernels/where_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/where_op.cc @@ -162,7 +162,8 @@ absl::StatusOr CompileWhereWithSort(XlaOpKernelContext* ctx) { TF_ASSIGN_OR_RETURN(xla::Shape input_shape, ctx->builder()->GetShape(condition)); auto iota_shape = - xla::ShapeUtil::MakeShape(xla::S32, input_shape.dimensions()); + xla::ShapeUtil::MakeShape(xla::S32, input_shape.dimensions(), + input_shape.expressions()); int64_t flattened_size = xla::Product(iota_shape.dimensions()); xla::DExpr flattened_expr = xla::DExpr::Const(1); @@ -192,7 +193,8 @@ absl::StatusOr CompileWhereWithSort(XlaOpKernelContext* ctx) { for (int64_t i = 0; i < iota_shape.dimensions_size(); ++i) { XlaOp index_single_dim = xla::GetTupleElement(sorted, i + 1); to_concat.push_back(xla::Reshape(index_single_dim, {flattened_size, 1}, - {flattened_expr, xla::DExpr::Const(1)})); + {flattened_expr, + xla::DExpr::Const(1)})); } XlaOp result = xla::ConcatInDim(ctx->builder(), to_concat, 1); @@ -264,8 +266,7 @@ absl::StatusOr CompileWhereWithPrefixSum(XlaOpKernelContext* ctx) { XlaOp out_idxs = xla::Select(xla::Ne(prefix_sum, prefix_sum_shifted), /*on_true=*/prefix_sum - xla::One(b, S32), /*on_false=*/oob_idx); - out_idxs = xla::Reshape(out_idxs, {flattened_size, 1}, - {flattened_expr, xla::DExpr::Const(1)}); + out_idxs = xla::Reshape(out_idxs, {flattened_size, 1}); // tf.where returns an array of multidimensional indices where the condition // is true. For example: @@ -288,12 +289,16 @@ absl::StatusOr CompileWhereWithPrefixSum(XlaOpKernelContext* ctx) { // // and then scatter iotas[out_idxs] into the output. std::vector iotas_to_concat; - auto iota_shape = xla::ShapeUtil::MakeShape(S32, input_shape.dimensions()); + auto iota_shape = xla::ShapeUtil::MakeShape( + S32, input_shape.dimensions(), input_shape.expressions()); iotas_to_concat.reserve(iota_shape.dimensions_size()); for (int64_t axis = 0; axis < iota_shape.dimensions_size(); ++axis) { - iotas_to_concat.push_back( - xla::Reshape(xla::Iota(b, iota_shape, axis), {flattened_size, 1}, - {flattened_expr, xla::DExpr::Const(1)})); + XlaOp flattened_iota = + xla::Reshape(xla::Iota(b, iota_shape, axis), {flattened_size}, + {flattened_expr}); + iotas_to_concat.push_back(xla::Reshape( + flattened_iota, {flattened_size, 1}, + {flattened_expr, xla::DExpr::Const(1)})); } XlaOp iotas = xla::ConcatInDim(b, iotas_to_concat, /*dimension=*/1); @@ -318,7 +323,12 @@ absl::StatusOr CompileWhereWithPrefixSum(XlaOpKernelContext* ctx) { XlaOp scattered = xla::Scatter( /*input=*/xla::Zeros( b, /*shape=*/xla::ShapeUtil::MakeShape( - S32, {flattened_size, iota_shape.dimensions_size()})), + S32, + std::vector{flattened_size, + iota_shape.dimensions_size()}, + std::vector{ + flattened_expr, + xla::DExpr::Const(iota_shape.dimensions_size())})), /*scatter_indices=*/out_idxs, /*updates=*/iotas, /*update_computation=*/assn_computation, scatter_dnums, /*indices_are_sorted=*/true, /*unique_indices=*/true); diff --git a/tensorflow/compiler/tf2xla/shape_util.cc b/tensorflow/compiler/tf2xla/shape_util.cc index 6aaa0419966e37..132daea9ee73e7 100644 --- a/tensorflow/compiler/tf2xla/shape_util.cc +++ b/tensorflow/compiler/tf2xla/shape_util.cc @@ -101,8 +101,7 @@ absl::Status XLAShapeToTensorShape(const xla::Shape& shape, for (int i = 0; i < shape.dimensions().size(); ++i) { TF_RETURN_IF_ERROR(tensor_shape->AddDimWithStatus(shape.dimensions(i))); } - MarkForCompilationPassFlags* flags = GetMarkForCompilationPassFlags(); - if (flags->tf_xla_enable_dynamic_sizes) { + if (!shape.expressions().empty()) { std::vector dexprs(shape.expressions().begin(), shape.expressions().end()); tensor_shape->set_expressions(std::move(dexprs)); diff --git a/tensorflow/compiler/tf2xla/xla_compiler.cc b/tensorflow/compiler/tf2xla/xla_compiler.cc index 219c51175b5743..90faea1745d327 100644 --- a/tensorflow/compiler/tf2xla/xla_compiler.cc +++ b/tensorflow/compiler/tf2xla/xla_compiler.cc @@ -946,14 +946,17 @@ absl::Status XlaCompiler::XLAShapeForArgument( TF_RETURN_IF_ERROR(RewriteLayoutWithShardedShape( arg_sharding, /*use_fast_memory=*/false, options_.shape_determination_fns, xla_shape)); - // If the arg is dynamic then we update the shape to reflect that. The - // layout etc above lose it by forcing a swap to TensorShape. - if (std::holds_alternative(arg.shape) && - std::get(arg.shape).is_dynamic()) { - xla::Shape dynamic_shape = std::get(arg.shape); - for (int i = 0; i < xla_shape->dimensions().size(); ++i) { - xla_shape->set_dynamic_dimension( - i, dynamic_shape.is_dynamic_dimension(i)); + // If the arg carries dynamic metadata or symbolic expressions then we + // update the shape to reflect that. The layout logic above routes + // through TensorShape and can otherwise discard this information. + if (std::holds_alternative(arg.shape)) { + const xla::Shape& original_shape = std::get(arg.shape); + if (original_shape.is_dynamic() || original_shape.has_dynamic_expr()) { + for (int i = 0; i < xla_shape->dimensions().size(); ++i) { + xla_shape->set_dynamic_dimension( + i, original_shape.is_dynamic_dimension(i)); + xla_shape->set_expression(i, original_shape.expressions(i)); + } } } } else { diff --git a/tensorflow/compiler/tf2xla/xla_compiler_test.cc b/tensorflow/compiler/tf2xla/xla_compiler_test.cc index 5aef5601af61ca..a955a88b2512a3 100644 --- a/tensorflow/compiler/tf2xla/xla_compiler_test.cc +++ b/tensorflow/compiler/tf2xla/xla_compiler_test.cc @@ -16,8 +16,10 @@ limitations under the License. #include "tensorflow/compiler/tf2xla/xla_compiler.h" #include +#include #include "absl/strings/match.h" #include "absl/strings/str_cat.h" +#include "tensorflow/compiler/jit/flags.h" #include "tensorflow/cc/framework/ops.h" #include "tensorflow/cc/ops/const_op.h" #include "tensorflow/cc/ops/data_flow_ops.h" @@ -96,6 +98,36 @@ class XlaCompilerTest : public ::testing::Test { std::unique_ptr flib_def_; }; +class ScopedTfXlaDynamicSizesFlag { + public: + ScopedTfXlaDynamicSizesFlag() { + old_value_ = GetMarkForCompilationPassFlags()->tf_xla_enable_dynamic_sizes; + GetMarkForCompilationPassFlags()->tf_xla_enable_dynamic_sizes = true; + SetTensorShapeExpressionsEnabledForTesting(true); + } + + ~ScopedTfXlaDynamicSizesFlag() { + SetTensorShapeExpressionsEnabledForTesting(std::nullopt); + GetMarkForCompilationPassFlags()->tf_xla_enable_dynamic_sizes = old_value_; + } + + private: + bool old_value_ = false; +}; + +class XlaCompilerDynamicSizesTest : public XlaCompilerTest { + protected: + void SetUp() override { + dynamic_sizes_flag_ = std::make_unique(); + XlaCompilerTest::SetUp(); + } + + void TearDown() override { dynamic_sizes_flag_.reset(); } + + private: + std::unique_ptr dynamic_sizes_flag_; +}; + namespace { // Helper class to test the ability to pass resources through to XLA @@ -106,193 +138,2713 @@ class DummyResourceForTest : public ResourceBase { void Increment() { ++value_; } int Get() { return value_; } - private: - int value_ = 0; -}; + private: + int value_ = 0; +}; + +class DummyReadResourceOp : public XlaOpKernel { + public: + explicit DummyReadResourceOp(OpKernelConstruction* ctx) : XlaOpKernel(ctx) {} + void Compile(XlaOpKernelContext* ctx) override { + ResourceMgr* rm = ctx->op_kernel_context()->resource_manager(); + OP_REQUIRES(ctx, rm, errors::Internal("No resource manager.")); + DummyResourceForTest* dummy; + OP_REQUIRES_OK(ctx, rm->Lookup( + rm->default_container(), "dummy", &dummy)); + dummy->Increment(); + dummy->Unref(); + + ctx->SetOutput(0, ctx->Input(0)); + ctx->SetOutput(1, ctx->Input(0)); + } +}; + +class DummyReadResourceCC { + public: + DummyReadResourceCC(const Scope& scope, const Input& value) { + if (!scope.ok()) return; + auto _value = ops::AsNodeOut(scope, value); + if (!scope.ok()) return; + Node* ret; + const auto unique_name = scope.GetUniqueNameForOp("DummyReadResource"); + auto builder = NodeBuilder(unique_name, "DummyReadResource").Input(_value); + scope.UpdateBuilder(&builder); + scope.UpdateStatus(builder.Finalize(scope.graph(), &ret)); + if (!scope.ok()) return; + scope.UpdateStatus(scope.DoShapeInference(ret)); + if (!scope.ok()) return; + this->output1_ = Output(ret, 0); + this->output2_ = Output(ret, 1); + } + + Output output1_; + Output output2_; +}; + +REGISTER_OP("DummyReadResource") + .Input("input: int32") + .Output("output1: int32") + .Output("output2: int32") + .SetShapeFn(shape_inference::UnknownShape) + .Doc(R"doc( +A dummy Op. + +input: dummy input. +output1: dummy output. +output2: dummy output. +)doc"); + +REGISTER_XLA_OP(Name("DummyReadResource"), DummyReadResourceOp); + +// DummyDuplicateOp is present purely to test multiple REGISTER_XLA_OP calls +// on the same Op name below. +class DummyDuplicateOp : public XlaOpKernel { + public: + explicit DummyDuplicateOp(OpKernelConstruction* ctx) : XlaOpKernel(ctx) {} + void Compile(XlaOpKernelContext* ctx) override { + ctx->SetOutput(0, ctx->Input(0)); + } +}; + +REGISTER_OP("DummyDuplicateOp") + .Input("input: int32") + .Output("output: int32") + .Doc(R"doc( +A dummy Op. + +input: dummy input. +output: dummy output. +)doc"); + +REGISTER_XLA_OP(Name("DummyDuplicateOp").Device(DEVICE_CPU_XLA_JIT), + DummyDuplicateOp); +REGISTER_XLA_OP(Name("DummyDuplicateOp").Device(DEVICE_GPU_XLA_JIT), + DummyDuplicateOp); + +// Tests compilation and execution of an empty graph. +TEST_F(XlaCompilerTest, EmptyReturnValues) { + XlaCompiler compiler(DefaultOptions()); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "add", + std::move(graph), + /*args=*/{}, &result)); + + TF_ASSERT_OK(client_->Execute(*result.computation, {}).status()); +} + +// Tests compilation and execution of a graph that adds two tensors. +TEST_F(XlaCompilerTest, Simple) { + // Builds a graph that adds two Tensors. + Scope scope = Scope::NewRootScope().ExitOnError(); + auto a = ops::_Arg(scope.WithOpName("A"), DT_INT32, 0); + auto b = ops::_Arg(scope.WithOpName("B"), DT_INT32, 1); + auto c = ops::Add(scope.WithOpName("C"), a, b); + auto d = ops::_Retval(scope.WithOpName("D"), c, 0); + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + // Builds a description of the arguments. + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = TensorShape({2}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = TensorShape({2}); + + // Compiles the graph. + XlaCompiler compiler(DefaultOptions()); + + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "add", + std::move(graph), args, &result)); + + // Tests that the generated computation works. + xla::Literal param0_literal = xla::LiteralUtil::CreateR1({7, 42}); + xla::Literal param1_literal = xla::LiteralUtil::CreateR1({-3, 101}); + std::unique_ptr param0_data = + client_->TransferToServer(param0_literal).value(); + std::unique_ptr param1_data = + client_->TransferToServer(param1_literal).value(); + + std::unique_ptr actual = + client_ + ->Execute(*result.computation, {param0_data.get(), param1_data.get()}) + .value(); + xla::Literal actual_literal = client_->Transfer(*actual).value(); + + xla::Literal expected0 = xla::LiteralUtil::CreateR1({4, 143}); + xla::Literal expected_literal = xla::LiteralUtil::MakeTuple({&expected0}); + EXPECT_TRUE(xla::LiteralTestUtil::Equal(expected_literal, actual_literal)); +} + +absl::StatusOr> LoadModuleFromHloProto( + const xla::HloModuleProto& module_proto) { + TF_ASSIGN_OR_RETURN(auto module_config, + xla::HloModule::CreateModuleConfigFromProto( + module_proto, xla::GetDebugOptionsFromFlags())); + return xla::CreateModuleFromProto(module_proto, module_config); +} + +// Tests compilation and execution of a graph that adds two tensors with dynamic +// shape parameters. +TEST_F(XlaCompilerTest, SimpleDynamicShapeParameter) { + // Builds a graph that adds two Tensors. + Scope scope = Scope::NewRootScope().ExitOnError(); + auto a = ops::_Arg(scope.WithOpName("A"), DT_INT32, 0); + auto b = ops::_Arg(scope.WithOpName("B"), DT_INT32, 1); + auto c = ops::Add(scope.WithOpName("C"), a, b); + auto d = ops::_Retval(scope.WithOpName("D"), c, 0); + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + // Builds a description of the arguments. + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = + xla::ShapeUtil::MakeShape(/*element_type=*/xla::S32, /*dimensions=*/{2}, + /*dynamic_dimensions=*/std::vector{true}, + /*expressions=*/{}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = TensorShape(/*dimensions=*/{2}); + + // Compiles the graph. + XlaCompiler compiler(DefaultOptions()); + + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "add", + std::move(graph), args, &result)); + + auto hlo = result.computation->proto(); + TF_ASSERT_OK_AND_ASSIGN(auto module, LoadModuleFromHloProto(hlo)); + EXPECT_EQ(module->computation_count(), 1); + EXPECT_TRUE(module->mutable_computation(0) + ->parameter_instruction(0) + ->shape() + .is_dynamic()); +} + +TEST_F(XlaCompilerDynamicSizesTest, DynamicShapeParameterPreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto identity = ops::Identity(scope.WithOpName("identity"), input); + auto retval = ops::_Retval(scope.WithOpName("retval"), identity, 0); + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6}, std::vector{xla::DExpr::Var(1)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "identity", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(1))); + + TF_ASSERT_OK_AND_ASSIGN(auto module, + LoadModuleFromHloProto(result.computation->proto())); + const xla::Shape& param_shape = + module->entry_computation()->parameter_instruction(0)->shape(); + EXPECT_TRUE( + xla::DynExpr::equal(param_shape.expressions(0), xla::DExpr::Var(1))); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(0), xla::DExpr::Var(1))); +} + +TEST_F(XlaCompilerDynamicSizesTest, ReverseSequencePreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto seq_lens = ops::_Arg(scope.WithOpName("seq_lens"), DT_INT32, 1); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("reverse_sequence", "ReverseSequence") + .Input(input.node()->name(), 0, DT_INT32) + .Input(seq_lens.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Tlen", DT_INT32) + .Attr("batch_dim", 0) + .Attr("seq_dim", 1) + .Finalize(&def)); + absl::Status status; + Node* reverse_sequence = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + TF_ASSERT_OK(scope.DoShapeInference(reverse_sequence)); + scope.graph()->AddEdge(input.node(), 0, reverse_sequence, 0); + scope.graph()->AddEdge(seq_lens.node(), 0, reverse_sequence, 1); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(reverse_sequence), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {4, 8}, + std::vector{xla::DExpr::Var(1), xla::DExpr::Var(2)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = TensorShape({4}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "reverse_sequence", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(1))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Var(2))); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(0), xla::DExpr::Var(1))); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(1), xla::DExpr::Var(2))); +} + +TEST_F(XlaCompilerDynamicSizesTest, UniquePreservesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("unique", "Unique") + .Input(input.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("out_idx", DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* unique = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + TF_ASSERT_OK(scope.DoShapeInference(unique)); + scope.graph()->AddEdge(input.node(), 0, unique, 0); + + auto retval0 = + ops::_Retval(scope.WithOpName("retval0"), Output(unique, 0), 0); + auto retval1 = + ops::_Retval(scope.WithOpName("retval1"), Output(unique, 1), 1); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7}, std::vector{xla::DExpr::Var(3)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "unique", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 2); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(3))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[1].shape.get_filled_expression( + 0), + xla::DExpr::Var(3))); + + const xla::Shape& values_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + const xla::Shape& indices_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {1}); + EXPECT_TRUE( + xla::DynExpr::equal(values_shape.expressions(0), xla::DExpr::Var(3))); + EXPECT_TRUE( + xla::DynExpr::equal(indices_shape.expressions(0), xla::DExpr::Var(3))); +} + +TEST_F(XlaCompilerDynamicSizesTest, DynamicPartitionPreservesPartitionExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto data = ops::_Arg(scope.WithOpName("data"), DT_INT32, 0); + auto partitions = ops::_Arg(scope.WithOpName("partitions"), DT_INT32, 1); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("dynamic_partition", "DynamicPartition") + .Input(data.node()->name(), 0, DT_INT32) + .Input(partitions.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("num_partitions", 2) + .Finalize(&def)); + absl::Status status; + Node* dynamic_partition = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + TF_ASSERT_OK(scope.DoShapeInference(dynamic_partition)); + scope.graph()->AddEdge(data.node(), 0, dynamic_partition, 0); + scope.graph()->AddEdge(partitions.node(), 0, dynamic_partition, 1); + + auto retval0 = ops::_Retval(scope.WithOpName("retval0"), + Output(dynamic_partition, 0), 0); + auto retval1 = ops::_Retval(scope.WithOpName("retval1"), + Output(dynamic_partition, 1), 1); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6}, std::vector{xla::DExpr::Var(4)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6}, std::vector{xla::DExpr::Var(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "dynamic_partition", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 2); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(4))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[1].shape.get_filled_expression( + 0), + xla::DExpr::Var(4))); + + const xla::Shape& result0_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + const xla::Shape& result1_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {1}); + EXPECT_TRUE( + xla::DynExpr::equal(result0_shape.expressions(0), xla::DExpr::Var(4))); + EXPECT_TRUE( + xla::DynExpr::equal(result1_shape.expressions(0), xla::DExpr::Var(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, + DynamicPartitionBroadcastPreservesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto data = ops::_Arg(scope.WithOpName("data"), DT_INT32, 0); + auto partitions = ops::_Arg(scope.WithOpName("partitions"), DT_INT32, 1); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("dynamic_partition", "DynamicPartition") + .Input(data.node()->name(), 0, DT_INT32) + .Input(partitions.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("num_partitions", 2) + .Finalize(&def)); + absl::Status status; + Node* dynamic_partition = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + TF_ASSERT_OK(scope.DoShapeInference(dynamic_partition)); + scope.graph()->AddEdge(data.node(), 0, dynamic_partition, 0); + scope.graph()->AddEdge(partitions.node(), 0, dynamic_partition, 1); + + auto retval0 = ops::_Retval(scope.WithOpName("retval0"), + Output(dynamic_partition, 0), 0); + auto retval1 = ops::_Retval(scope.WithOpName("retval1"), + Output(dynamic_partition, 1), 1); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 3}, + std::vector{xla::DExpr::Var(40), xla::DExpr::Const(3)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8}, std::vector{xla::DExpr::Var(40)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "dynamic_partition_broadcast", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 2); + for (int i = 0; i < 2; ++i) { + EXPECT_TRUE(xla::DynExpr::equal( + result.outputs[i].shape.get_filled_expression(0), xla::DExpr::Var(40))); + EXPECT_TRUE(xla::DynExpr::equal( + result.outputs[i].shape.get_filled_expression(1), xla::DExpr::Const(3))); + const xla::Shape& out_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {i}); + EXPECT_TRUE( + xla::DynExpr::equal(out_shape.expressions(0), xla::DExpr::Var(40))); + EXPECT_TRUE( + xla::DynExpr::equal(out_shape.expressions(1), xla::DExpr::Const(3))); + } +} + +TEST_F(XlaCompilerDynamicSizesTest, + DenseBincountMatrixPreservesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto weights = ops::_Arg(scope.WithOpName("weights"), DT_FLOAT, 1); + auto size = ops::Const(scope.WithOpName("size"), 5); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("dense_bincount", "DenseBincount") + .Input(input.node()->name(), 0, DT_INT32) + .Input(size.node()->name(), 0, DT_INT32) + .Input(weights.node()->name(), 0, DT_FLOAT) + .Attr("Tidx", DT_INT32) + .Attr("T", DT_FLOAT) + .Attr("binary_output", false) + .Finalize(&def)); + absl::Status status; + Node* bincount = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, bincount, 0); + scope.graph()->AddEdge(size.node(), 0, bincount, 1); + scope.graph()->AddEdge(weights.node(), 0, bincount, 2); + TF_ASSERT_OK(scope.DoShapeInference(bincount)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(bincount, 0), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {12, 4}, + std::vector{xla::DExpr::Var(50), xla::DExpr::Const(4)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_FLOAT; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::F32, {12, 4}, + std::vector{xla::DExpr::Var(50), xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "dense_bincount", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(50))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); + + const xla::Shape& out_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE(xla::DynExpr::equal(out_shape.expressions(0), xla::DExpr::Var(50))); + EXPECT_TRUE( + xla::DynExpr::equal(out_shape.expressions(1), xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, + DynamicStitchEmptyPreservesTrailingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto data0 = ops::_Arg(scope.WithOpName("data0"), DT_INT32, 0); + auto data1 = ops::_Arg(scope.WithOpName("data1"), DT_INT32, 1); + Tensor empty_indices_tensor(DT_INT32, TensorShape({0})); + auto indices0 = ops::Const(scope.WithOpName("indices0"), empty_indices_tensor); + auto indices1 = ops::Const(scope.WithOpName("indices1"), empty_indices_tensor); + + NodeDef def; + std::vector indices_inputs = { + {indices0.node()->name(), 0, DT_INT32}, + {indices1.node()->name(), 0, DT_INT32}, + }; + std::vector data_inputs = { + {data0.node()->name(), 0, DT_INT32}, + {data1.node()->name(), 0, DT_INT32}, + }; + TF_ASSERT_OK(NodeDefBuilder("dynamic_stitch_empty", "DynamicStitch") + .Input(indices_inputs) + .Input(data_inputs) + .Attr("N", 2) + .Attr("T", DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* dynamic_stitch = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(indices0.node(), 0, dynamic_stitch, 0); + scope.graph()->AddEdge(indices1.node(), 0, dynamic_stitch, 1); + scope.graph()->AddEdge(data0.node(), 0, dynamic_stitch, 2); + scope.graph()->AddEdge(data1.node(), 0, dynamic_stitch, 3); + + auto retval = ops::_Retval(scope.WithOpName("retval"), + Output(dynamic_stitch, 0), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {0, 7}, + std::vector{xla::DExpr::Const(0), xla::DExpr::Var(61)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {0, 7}, + std::vector{xla::DExpr::Const(0), xla::DExpr::Var(61)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "dynamic_stitch_empty", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Const(0))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Var(61))); + + const xla::Shape& out_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(out_shape.expressions(0), xla::DExpr::Const(0))); + EXPECT_TRUE( + xla::DynExpr::equal(out_shape.expressions(1), xla::DExpr::Var(61))); +} + +TEST_F(XlaCompilerDynamicSizesTest, + TensorListPushBackStackPreservesElementExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto element = ops::_Arg(scope.WithOpName("element"), DT_INT32, 0); + auto element_shape = ops::Const(scope.WithOpName("element_shape"), + {-1, 3}, {2}); + auto max_num_elements = ops::Const(scope.WithOpName("max_num_elements"), 4); + + NodeDef empty_def; + TF_ASSERT_OK(NodeDefBuilder("empty_list", "EmptyTensorList") + .Input(element_shape.node()->name(), 0, DT_INT32) + .Input(max_num_elements.node()->name(), 0, DT_INT32) + .Attr("element_dtype", DT_INT32) + .Attr("shape_type", DT_INT32) + .Finalize(&empty_def)); + absl::Status status; + Node* empty_list = scope.graph()->AddNode(empty_def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(element_shape.node(), 0, empty_list, 0); + scope.graph()->AddEdge(max_num_elements.node(), 0, empty_list, 1); + TF_ASSERT_OK(scope.DoShapeInference(empty_list)); + + NodeDef push_def; + TF_ASSERT_OK(NodeDefBuilder("push_back", "TensorListPushBack") + .Input(empty_list->name(), 0, DT_VARIANT) + .Input(element.node()->name(), 0, DT_INT32) + .Attr("element_dtype", DT_INT32) + .Finalize(&push_def)); + Node* push_back = scope.graph()->AddNode(push_def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(empty_list, 0, push_back, 0); + scope.graph()->AddEdge(element.node(), 0, push_back, 1); + TF_ASSERT_OK(scope.DoShapeInference(push_back)); + + NodeDef stack_def; + TF_ASSERT_OK(NodeDefBuilder("stack", "TensorListStack") + .Input(push_back->name(), 0, DT_VARIANT) + .Input(element_shape.node()->name(), 0, DT_INT32) + .Attr("element_dtype", DT_INT32) + .Attr("num_elements", 4) + .Finalize(&stack_def)); + Node* stack = scope.graph()->AddNode(stack_def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(push_back, 0, stack, 0); + scope.graph()->AddEdge(element_shape.node(), 0, stack, 1); + TF_ASSERT_OK(scope.DoShapeInference(stack)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(stack, 0), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {5, 3}, + std::vector{xla::DExpr::Var(70), xla::DExpr::Const(3)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "tensor_list_stack", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Const(4))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Var(70))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(3))); +} + +TEST_F(XlaCompilerDynamicSizesTest, ShapeThenReshapePreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto shape_source = ops::_Arg(scope.WithOpName("shape_source"), DT_INT32, 0); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 1); + auto shape = ops::Shape(scope.WithOpName("shape"), shape_source); + auto reshaped = ops::Reshape(scope.WithOpName("reshape"), input, shape); + auto retval = ops::_Retval(scope.WithOpName("retval"), reshaped, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6, 7}, + std::vector{xla::DExpr::Var(41), xla::DExpr::Const(7)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6, 7}, + std::vector{xla::DExpr::Var(41), xla::DExpr::Const(7)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "shape_then_reshape", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(41))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(7))); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(0), xla::DExpr::Var(41))); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(1), xla::DExpr::Const(7))); +} + +TEST_F(XlaCompilerDynamicSizesTest, ZerosLikePreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto zeros = ops::ZerosLike(scope.WithOpName("zeros_like"), input); + auto retval = ops::_Retval(scope.WithOpName("retval"), zeros, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {9, 4}, + std::vector{xla::DExpr::Var(43), xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "zeros_like_exprs", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(43))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, OnesLikePreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto ones = ops::OnesLike(scope.WithOpName("ones_like"), input); + auto retval = ops::_Retval(scope.WithOpName("retval"), ones, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {9, 4}, + std::vector{xla::DExpr::Var(44), xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "ones_like_exprs", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(44))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, MatrixDiagPreservesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("matrix_diag", "MatrixDiag") + .Input(input.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* matrix_diag = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + TF_ASSERT_OK(scope.DoShapeInference(matrix_diag)); + scope.graph()->AddEdge(input.node(), 0, matrix_diag, 0); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(matrix_diag), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6, 4}, + std::vector{xla::DExpr::Var(45), xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "matrix_diag_exprs", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(45))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, WhereBuildsDynamicIndexMatrixShape) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_BOOL, 0); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("where", "Where") + .Input(input.node()->name(), 0, DT_BOOL) + .Finalize(&def)); + absl::Status status; + Node* where = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + TF_ASSERT_OK(scope.DoShapeInference(where)); + scope.graph()->AddEdge(input.node(), 0, where, 0); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(where), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_BOOL; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::PRED, {8, 4, 6}, + std::vector{xla::DExpr::Var(46), xla::DExpr::Const(4), + xla::DExpr::Var(47)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "where", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_EQ(result_shape.dimensions_size(), 2); + EXPECT_EQ(result_shape.dimensions(1), 3); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(1), xla::DExpr::Const(3))); +} + +TEST_F(XlaCompilerDynamicSizesTest, DiagDuplicatesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("diag", "Diag") + .Input(input.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* diag = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + TF_ASSERT_OK(scope.DoShapeInference(diag)); + scope.graph()->AddEdge(input.node(), 0, diag, 0); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(diag), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {5}, std::vector{xla::DExpr::Var(42)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "diag", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(42))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Var(42))); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(0), xla::DExpr::Var(42))); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(1), xla::DExpr::Var(42))); +} + +TEST_F(XlaCompilerDynamicSizesTest, InTopKPreservesBatchExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto predictions = ops::_Arg(scope.WithOpName("predictions"), DT_FLOAT, 0); + auto targets = ops::_Arg(scope.WithOpName("targets"), DT_INT32, 1); + auto k = ops::Const(scope.WithOpName("k"), 3, {}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("in_topk", "InTopKV2") + .Input(predictions.node()->name(), 0, DT_FLOAT) + .Input(targets.node()->name(), 0, DT_INT32) + .Input(k.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* in_topk = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(predictions.node(), 0, in_topk, 0); + scope.graph()->AddEdge(targets.node(), 0, in_topk, 1); + scope.graph()->AddEdge(k.node(), 0, in_topk, 2); + TF_ASSERT_OK(scope.DoShapeInference(in_topk)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(in_topk), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_FLOAT; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::F32, {5, 7}, + std::vector{xla::DExpr::Var(43), xla::DExpr::Const(7)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {5}, std::vector{xla::DExpr::Var(43)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "in_topk", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(43))); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(0), xla::DExpr::Var(43))); +} + +TEST_F(XlaCompilerDynamicSizesTest, ReshapeCollapsePreservesSymbolicExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto shape = ops::Const(scope.WithOpName("shape"), {96}, {1}); + auto reshaped = ops::Reshape(scope.WithOpName("reshape"), input, shape); + auto retval = ops::_Retval(scope.WithOpName("retval"), reshaped, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {3, 4, 8}, + std::vector{xla::DExpr::Var(5), xla::DExpr::Const(4), + xla::DExpr::Const(8)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "reshape_collapse", std::move(graph), args, + &result)); + + xla::DExpr expected = + (xla::DExpr::Var(5) * xla::DExpr::Const(32)).simplify(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + expected)); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE(xla::DynExpr::equal(result_shape.expressions(0), expected)); +} + +TEST_F(XlaCompilerDynamicSizesTest, ReshapeSplitPreservesSymbolicExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto shape = ops::Const(scope.WithOpName("shape"), {5, 16}, {2}); + auto reshaped = ops::Reshape(scope.WithOpName("reshape"), input, shape); + auto retval = ops::_Retval(scope.WithOpName("retval"), reshaped, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {10, 8}, + std::vector{xla::DExpr::Var(6), xla::DExpr::Const(8)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "reshape_split", std::move(graph), args, + &result)); + + xla::DExpr expected = + (xla::DExpr::Var(6) / xla::DExpr::Const(2)).simplify(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + expected)); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE(xla::DynExpr::equal(result_shape.expressions(0), expected)); +} + +TEST_F(XlaCompilerDynamicSizesTest, ReshapeSplitAndCollapsePreservesSymbolicExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto shape = ops::Const(scope.WithOpName("shape"), {4, 64}, {2}); + auto reshaped = ops::Reshape(scope.WithOpName("reshape"), input, shape); + auto retval = ops::_Retval(scope.WithOpName("retval"), reshaped, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 8, 4}, + std::vector{xla::DExpr::Var(7), xla::DExpr::Const(8), + xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "reshape_split_collapse", + std::move(graph), args, &result)); + + xla::DExpr expected = + (xla::DExpr::Var(7) / xla::DExpr::Const(2)).simplify(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + expected)); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE(xla::DynExpr::equal(result_shape.expressions(0), expected)); +} + +TEST_F(XlaCompilerDynamicSizesTest, GatherV2PreservesUngatheredExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto params = ops::_Arg(scope.WithOpName("params"), DT_INT32, 0); + auto indices = ops::Const(scope.WithOpName("indices"), {0, 2, 4}, {3}); + auto axis = ops::Const(scope.WithOpName("axis"), 1, {}); + auto gathered = + ops::GatherV2(scope.WithOpName("gather"), params, indices, axis); + auto retval = ops::_Retval(scope.WithOpName("retval"), gathered, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {9, 7}, + std::vector{xla::DExpr::Var(8), xla::DExpr::Const(7)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "gather_preserve", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(8))); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(0), xla::DExpr::Var(8))); +} + +TEST_F(XlaCompilerDynamicSizesTest, TransposePermutesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto perm = ops::Const(scope.WithOpName("perm"), {1, 0}, {2}); + auto transposed = ops::Transpose(scope.WithOpName("transpose"), input, perm); + auto retval = ops::_Retval(scope.WithOpName("retval"), transposed, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {5, 7}, + std::vector{xla::DExpr::Var(9), xla::DExpr::Var(10)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "transpose_exprs", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(10))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Var(9))); +} + +TEST_F(XlaCompilerDynamicSizesTest, ExpandDimsInsertsUnitExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto dim = ops::Const(scope.WithOpName("dim"), 1, {}); + auto expanded = ops::ExpandDims(scope.WithOpName("expand"), input, dim); + auto retval = ops::_Retval(scope.WithOpName("retval"), expanded, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6, 4}, + std::vector{xla::DExpr::Var(11), xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "expand_dims_exprs", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(11))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(1))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, SqueezeRemovesUnitExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("squeeze", "Squeeze") + .Input(input.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("squeeze_dims", {1}) + .Finalize(&def)); + absl::Status status; + Node* squeeze = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, squeeze, 0); + TF_ASSERT_OK(scope.DoShapeInference(squeeze)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(squeeze), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6, 1, 4}, + std::vector{xla::DExpr::Var(12), xla::DExpr::Const(1), + xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "squeeze_exprs", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(12))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, SplitPreservesAndDividesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto split_dim = ops::Const(scope.WithOpName("split_dim"), 0, {}); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto split = ops::Split(scope.WithOpName("split"), split_dim, input, 2); + auto retval0 = ops::_Retval(scope.WithOpName("retval0"), split.output[0], 0); + auto retval1 = ops::_Retval(scope.WithOpName("retval1"), split.output[1], 1); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 5}, + std::vector{xla::DExpr::Var(13), xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "split_exprs", std::move(graph), args, + &result)); + + xla::DExpr expected = + (xla::DExpr::Var(13) / xla::DExpr::Const(2)).simplify(); + ASSERT_EQ(result.outputs.size(), 2); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + expected)); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[1].shape.get_filled_expression( + 0), + expected)); +} + +TEST_F(XlaCompilerDynamicSizesTest, TileScalesExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto multiples = + ops::Const(scope.WithOpName("multiples"), {3, 1}, {2}); + auto tiled = ops::Tile(scope.WithOpName("tile"), input, multiples); + auto retval = ops::_Retval(scope.WithOpName("retval"), tiled, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {4, 5}, + std::vector{xla::DExpr::Var(14), xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "tile_exprs", std::move(graph), args, + &result)); + + xla::DExpr expected = + (xla::DExpr::Var(14) * xla::DExpr::Const(3)).simplify(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + expected)); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, PackInsertsAxisAndPreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input0 = ops::_Arg(scope.WithOpName("input0"), DT_INT32, 0); + auto input1 = ops::_Arg(scope.WithOpName("input1"), DT_INT32, 1); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("pack", "Pack") + .Input({NodeDefBuilder::NodeOut(input0.node()->name(), 0, + DT_INT32), + NodeDefBuilder::NodeOut(input1.node()->name(), 0, + DT_INT32)}) + .Attr("T", DT_INT32) + .Attr("N", 2) + .Attr("axis", 1) + .Finalize(&def)); + absl::Status status; + Node* pack = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input0.node(), 0, pack, 0); + scope.graph()->AddEdge(input1.node(), 0, pack, 1); + TF_ASSERT_OK(scope.DoShapeInference(pack)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(pack), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6}, std::vector{xla::DExpr::Var(15)}); + args[1] = args[0]; + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "pack", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(15))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(2))); +} + +TEST_F(XlaCompilerDynamicSizesTest, UnpackRemovesAxisAndPreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("unpack", "Unpack") + .Input(input.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("num", 3) + .Attr("axis", 2) + .Finalize(&def)); + absl::Status status; + Node* unpack = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, unpack, 0); + TF_ASSERT_OK(scope.DoShapeInference(unpack)); + + auto retval0 = ops::_Retval(scope.WithOpName("retval0"), Output(unpack, 0), 0); + auto retval1 = ops::_Retval(scope.WithOpName("retval1"), Output(unpack, 1), 1); + auto retval2 = ops::_Retval(scope.WithOpName("retval2"), Output(unpack, 2), 2); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 4, 3}, + std::vector{xla::DExpr::Var(16), xla::DExpr::Const(4), + xla::DExpr::Const(3)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "unpack", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 3); + for (int i = 0; i < 3; ++i) { + EXPECT_TRUE(xla::DynExpr::equal( + result.outputs[i].shape.get_filled_expression(0), xla::DExpr::Var(16))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[i].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); + } +} + +TEST_F(XlaCompilerDynamicSizesTest, ConcatV2AddsLeadingExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_INT32, 0); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_INT32, 1); + auto axis = ops::Const(scope.WithOpName("axis"), 0, {}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("concat", "ConcatV2") + .Input({NodeDefBuilder::NodeOut(lhs.node()->name(), 0, + DT_INT32), + NodeDefBuilder::NodeOut(rhs.node()->name(), 0, + DT_INT32)}) + .Input(axis.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Tidx", DT_INT32) + .Attr("N", 2) + .Finalize(&def)); + absl::Status status; + Node* concat = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(lhs.node(), 0, concat, 0); + scope.graph()->AddEdge(rhs.node(), 0, concat, 1); + scope.graph()->AddEdge(axis.node(), 0, concat, 2); + TF_ASSERT_OK(scope.DoShapeInference(concat)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(concat), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {5, 4}, + std::vector{xla::DExpr::Var(17), xla::DExpr::Const(4)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6, 4}, + std::vector{xla::DExpr::Var(18), xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "concat", + std::move(graph), args, &result)); + + xla::DExpr expected = + (xla::DExpr::Var(17) + xla::DExpr::Var(18)).simplify(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + expected)); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, ConcatAddsLeadingExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto axis = ops::Const(scope.WithOpName("axis"), 0, {}); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_INT32, 0); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_INT32, 1); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("concat", "Concat") + .Input(axis.node()->name(), 0, DT_INT32) + .Input({NodeDefBuilder::NodeOut(lhs.node()->name(), 0, + DT_INT32), + NodeDefBuilder::NodeOut(rhs.node()->name(), 0, + DT_INT32)}) + .Attr("T", DT_INT32) + .Attr("N", 2) + .Finalize(&def)); + absl::Status status; + Node* concat = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(axis.node(), 0, concat, 0); + scope.graph()->AddEdge(lhs.node(), 0, concat, 1); + scope.graph()->AddEdge(rhs.node(), 0, concat, 2); + TF_ASSERT_OK(scope.DoShapeInference(concat)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(concat), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {5, 4}, + std::vector{xla::DExpr::Var(22), xla::DExpr::Const(4)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6, 4}, + std::vector{xla::DExpr::Var(23), xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "concat_legacy", std::move(graph), args, + &result)); + + xla::DExpr expected = + (xla::DExpr::Var(22) + xla::DExpr::Var(23)).simplify(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + expected)); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, ConcatAddsThreeLeadingExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_INT32, 0); + auto mid = ops::_Arg(scope.WithOpName("mid"), DT_INT32, 1); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_INT32, 2); + auto axis = ops::Const(scope.WithOpName("axis"), 0, {}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("concat0", "ConcatV2") + .Input({NodeDefBuilder::NodeOut(lhs.node()->name(), 0, + DT_INT32), + NodeDefBuilder::NodeOut(mid.node()->name(), 0, + DT_INT32)}) + .Input(axis.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Tidx", DT_INT32) + .Attr("N", 2) + .Finalize(&def)); + absl::Status status; + Node* concat0 = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(lhs.node(), 0, concat0, 0); + scope.graph()->AddEdge(mid.node(), 0, concat0, 1); + scope.graph()->AddEdge(axis.node(), 0, concat0, 2); + TF_ASSERT_OK(scope.DoShapeInference(concat0)); + + NodeDef def1; + TF_ASSERT_OK(NodeDefBuilder("concat1", "ConcatV2") + .Input({NodeDefBuilder::NodeOut(concat0->name(), 0, + DT_INT32), + NodeDefBuilder::NodeOut(rhs.node()->name(), 0, + DT_INT32)}) + .Input(axis.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Tidx", DT_INT32) + .Attr("N", 2) + .Finalize(&def1)); + Node* concat1 = scope.graph()->AddNode(def1, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(concat0, 0, concat1, 0); + scope.graph()->AddEdge(rhs.node(), 0, concat1, 1); + scope.graph()->AddEdge(axis.node(), 0, concat1, 2); + TF_ASSERT_OK(scope.DoShapeInference(concat1)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(concat1), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(3); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {3, 4}, + std::vector{xla::DExpr::Var(31), xla::DExpr::Const(4)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {5, 4}, + std::vector{xla::DExpr::Var(32), xla::DExpr::Const(4)}); + args[2].kind = XlaCompiler::Argument::kParameter; + args[2].type = DT_INT32; + args[2].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 4}, + std::vector{xla::DExpr::Var(33), xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "concat_three", std::move(graph), args, + &result)); + + xla::DExpr expected = + (xla::DExpr::Var(31) + xla::DExpr::Var(32) + xla::DExpr::Var(33)) + .simplify(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + expected)); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE(xla::DynExpr::equal(result_shape.expressions(0), expected)); + EXPECT_TRUE(xla::DynExpr::equal(result_shape.expressions(1), + xla::DExpr::Const(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, ConcatV2PreservesLeadingExpressionOnInnerAxis) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_INT32, 0); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_INT32, 1); + auto axis = ops::Const(scope.WithOpName("axis"), 1, {}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("concat", "ConcatV2") + .Input({NodeDefBuilder::NodeOut(lhs.node()->name(), 0, + DT_INT32), + NodeDefBuilder::NodeOut(rhs.node()->name(), 0, + DT_INT32)}) + .Input(axis.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Tidx", DT_INT32) + .Attr("N", 2) + .Finalize(&def)); + absl::Status status; + Node* concat = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(lhs.node(), 0, concat, 0); + scope.graph()->AddEdge(rhs.node(), 0, concat, 1); + scope.graph()->AddEdge(axis.node(), 0, concat, 2); + TF_ASSERT_OK(scope.DoShapeInference(concat)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(concat), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 4}, + std::vector{xla::DExpr::Var(24), xla::DExpr::Const(4)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 3}, + std::vector{xla::DExpr::Var(24), xla::DExpr::Const(3)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "concat_inner", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(24))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(7))); +} + +TEST_F(XlaCompilerDynamicSizesTest, AddPreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_INT32, 0); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_INT32, 1); + auto sum = ops::Add(scope.WithOpName("add"), lhs, rhs); + auto retval = ops::_Retval(scope.WithOpName("retval"), sum, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 5}, + std::vector{xla::DExpr::Var(44), xla::DExpr::Const(5)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 5}, + std::vector{xla::DExpr::Var(44), xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "add", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(44))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, + AddSameRankBroadcastPreservesMappedExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_INT32, 0); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_INT32, 1); + auto sum = ops::Add(scope.WithOpName("add"), lhs, rhs); + auto retval = ops::_Retval(scope.WithOpName("retval"), sum, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 1, 3}, + std::vector{xla::DExpr::Var(45), xla::DExpr::Const(1), + xla::DExpr::Const(3)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 4, 3}, + std::vector{xla::DExpr::Var(45), xla::DExpr::Const(4), + xla::DExpr::Const(3)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "add_broadcast", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(45))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(3))); +} + +TEST_F(XlaCompilerDynamicSizesTest, AddDegenerateBroadcastPreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_INT32, 0); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_INT32, 1); + auto sum = ops::Add(scope.WithOpName("add"), lhs, rhs); + auto retval = ops::_Retval(scope.WithOpName("retval"), sum, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {1, 5}, + std::vector{xla::DExpr::Const(1), xla::DExpr::Const(5)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 5}, + std::vector{xla::DExpr::Var(46), xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "add_degenerate", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(46))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, + MulSameRankBroadcastPreservesMappedExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_INT32, 0); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_INT32, 1); + auto product = ops::Mul(scope.WithOpName("mul"), lhs, rhs); + auto retval = ops::_Retval(scope.WithOpName("retval"), product, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 1, 3}, + std::vector{xla::DExpr::Var(47), xla::DExpr::Const(1), + xla::DExpr::Const(3)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_INT32; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 4, 3}, + std::vector{xla::DExpr::Var(47), xla::DExpr::Const(4), + xla::DExpr::Const(3)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "mul_broadcast", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(47))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(3))); +} + +TEST_F(XlaCompilerDynamicSizesTest, ReverseV2PreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto axis = ops::Const(scope.WithOpName("axis"), {1}, {1}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("reverse", "ReverseV2") + .Input(input.node()->name(), 0, DT_INT32) + .Input(axis.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Tidx", DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* reverse = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, reverse, 0); + scope.graph()->AddEdge(axis.node(), 0, reverse, 1); + TF_ASSERT_OK(scope.DoShapeInference(reverse)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(reverse), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {8, 5}, + std::vector{xla::DExpr::Var(19), xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "reverse_v2", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(19))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, BatchMatMulPreservesBatchExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_FLOAT, 0); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_FLOAT, 1); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("batch_matmul", "BatchMatMul") + .Input(lhs.node()->name(), 0, DT_FLOAT) + .Input(rhs.node()->name(), 0, DT_FLOAT) + .Attr("T", DT_FLOAT) + .Attr("adj_x", false) + .Attr("adj_y", false) + .Finalize(&def)); + absl::Status status; + Node* batch_matmul = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(lhs.node(), 0, batch_matmul, 0); + scope.graph()->AddEdge(rhs.node(), 0, batch_matmul, 1); + TF_ASSERT_OK(scope.DoShapeInference(batch_matmul)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(batch_matmul), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_FLOAT; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::F32, {8, 4, 6}, + std::vector{xla::DExpr::Var(25), xla::DExpr::Const(4), + xla::DExpr::Const(6)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_FLOAT; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::F32, {8, 6, 5}, + std::vector{xla::DExpr::Var(25), xla::DExpr::Const(6), + xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "batch_matmul", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(25))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, BatchMatMulV2BroadcastsBatchExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto lhs = ops::_Arg(scope.WithOpName("lhs"), DT_FLOAT, 0); + auto rhs = ops::_Arg(scope.WithOpName("rhs"), DT_FLOAT, 1); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("batch_matmul_v2", "BatchMatMulV2") + .Input(lhs.node()->name(), 0, DT_FLOAT) + .Input(rhs.node()->name(), 0, DT_FLOAT) + .Attr("T", DT_FLOAT) + .Attr("adj_x", false) + .Attr("adj_y", false) + .Finalize(&def)); + absl::Status status; + Node* batch_matmul = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(lhs.node(), 0, batch_matmul, 0); + scope.graph()->AddEdge(rhs.node(), 0, batch_matmul, 1); + TF_ASSERT_OK(scope.DoShapeInference(batch_matmul)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(batch_matmul), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_FLOAT; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::F32, {1, 4, 6}, + std::vector{xla::DExpr::Const(1), xla::DExpr::Const(4), + xla::DExpr::Const(6)}); + args[1].kind = XlaCompiler::Argument::kParameter; + args[1].type = DT_FLOAT; + args[1].shape = xla::ShapeUtil::MakeShape( + xla::F32, {8, 6, 5}, + std::vector{xla::DExpr::Var(26), xla::DExpr::Const(6), + xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "batch_matmul_v2", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(26))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, SlicePreservesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto begin = ops::Const(scope.WithOpName("begin"), {0, 2}, {2}); + auto size = ops::Const(scope.WithOpName("size"), {-1, 3}, {2}); + auto sliced = ops::Slice(scope.WithOpName("slice"), input, begin, size); + auto retval = ops::_Retval(scope.WithOpName("retval"), sliced, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 8}, + std::vector{xla::DExpr::Var(20), xla::DExpr::Const(8)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "slice", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(20))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(3))); +} + +TEST_F(XlaCompilerDynamicSizesTest, SliceSubtractsFromLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto begin = ops::Const(scope.WithOpName("begin"), {2, 1}, {2}); + auto size = ops::Const(scope.WithOpName("size"), {-1, 3}, {2}); + auto sliced = ops::Slice(scope.WithOpName("slice"), input, begin, size); + auto retval = ops::_Retval(scope.WithOpName("retval"), sliced, 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {9, 8}, + std::vector{xla::DExpr::Var(27), xla::DExpr::Const(8)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "slice_subtract", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + (xla::DExpr::Var(27) - 2).simplify())); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(3))); +} + +TEST_F(XlaCompilerDynamicSizesTest, + StridedSlicePreservesLeadingExpressionOnInnerAxis) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto begin = ops::Const(scope.WithOpName("begin"), {0, 1}, {2}); + auto end = ops::Const(scope.WithOpName("end"), {0, 7}, {2}); + auto strides = ops::Const(scope.WithOpName("strides"), {1, 2}, {2}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("strided_slice", "StridedSlice") + .Input(input.node()->name(), 0, DT_INT32) + .Input(begin.node()->name(), 0, DT_INT32) + .Input(end.node()->name(), 0, DT_INT32) + .Input(strides.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Index", DT_INT32) + .Attr("begin_mask", 1) + .Attr("end_mask", 1) + .Attr("ellipsis_mask", 0) + .Attr("new_axis_mask", 0) + .Attr("shrink_axis_mask", 0) + .Finalize(&def)); + absl::Status status; + Node* strided_slice = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, strided_slice, 0); + scope.graph()->AddEdge(begin.node(), 0, strided_slice, 1); + scope.graph()->AddEdge(end.node(), 0, strided_slice, 2); + scope.graph()->AddEdge(strides.node(), 0, strided_slice, 3); + TF_ASSERT_OK(scope.DoShapeInference(strided_slice)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(strided_slice), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 8}, + std::vector{xla::DExpr::Var(40), xla::DExpr::Const(8)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "strided_slice_inner", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(40))); + EXPECT_EQ(result.outputs[0].shape.dim_size(1), 3); +} + +TEST_F(XlaCompilerDynamicSizesTest, StridedSliceScalesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto begin = ops::Const(scope.WithOpName("begin"), {0, 0}, {2}); + auto end = ops::Const(scope.WithOpName("end"), {0, 4}, {2}); + auto strides = ops::Const(scope.WithOpName("strides"), {2, 1}, {2}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("strided_slice", "StridedSlice") + .Input(input.node()->name(), 0, DT_INT32) + .Input(begin.node()->name(), 0, DT_INT32) + .Input(end.node()->name(), 0, DT_INT32) + .Input(strides.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Index", DT_INT32) + .Attr("begin_mask", 1) + .Attr("end_mask", 1) + .Attr("ellipsis_mask", 0) + .Attr("new_axis_mask", 0) + .Attr("shrink_axis_mask", 0) + .Finalize(&def)); + absl::Status status; + Node* strided_slice = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, strided_slice, 0); + scope.graph()->AddEdge(begin.node(), 0, strided_slice, 1); + scope.graph()->AddEdge(end.node(), 0, strided_slice, 2); + scope.graph()->AddEdge(strides.node(), 0, strided_slice, 3); + TF_ASSERT_OK(scope.DoShapeInference(strided_slice)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(strided_slice), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 4}, + std::vector{ + ((xla::DExpr::Const(2) * xla::DExpr::Var(41)) - + xla::DExpr::Const(1)) + .simplify(), + xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "strided_slice_leading", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal( + result.outputs[0].shape.get_filled_expression(0), xla::DExpr::Var(41))); + EXPECT_EQ(result.outputs[0].shape.dim_size(0), 4); + EXPECT_EQ(result.outputs[0].shape.dim_size(1), 4); +} + +TEST_F(XlaCompilerDynamicSizesTest, StridedSliceNewAxisInsertsUnitExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto begin = ops::Const(scope.WithOpName("begin"), {0, 0}, {2}); + auto end = ops::Const(scope.WithOpName("end"), {0, 0}, {2}); + auto strides = ops::Const(scope.WithOpName("strides"), {1, 1}, {2}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("strided_slice", "StridedSlice") + .Input(input.node()->name(), 0, DT_INT32) + .Input(begin.node()->name(), 0, DT_INT32) + .Input(end.node()->name(), 0, DT_INT32) + .Input(strides.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Index", DT_INT32) + .Attr("begin_mask", 0x1) + .Attr("end_mask", 0x1) + .Attr("ellipsis_mask", 0) + .Attr("new_axis_mask", 0x2) + .Attr("shrink_axis_mask", 0) + .Finalize(&def)); + absl::Status status; + Node* strided_slice = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, strided_slice, 0); + scope.graph()->AddEdge(begin.node(), 0, strided_slice, 1); + scope.graph()->AddEdge(end.node(), 0, strided_slice, 2); + scope.graph()->AddEdge(strides.node(), 0, strided_slice, 3); + TF_ASSERT_OK(scope.DoShapeInference(strided_slice)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(strided_slice), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 5}, + std::vector{xla::DExpr::Var(42), xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "strided_slice_new_axis", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(42))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(1))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, + StridedSliceShrinkAxisPreservesRemainingExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto begin = ops::Const(scope.WithOpName("begin"), {0, 2, 0}, {3}); + auto end = ops::Const(scope.WithOpName("end"), {0, 2, 5}, {3}); + auto strides = ops::Const(scope.WithOpName("strides"), {1, 1, 1}, {3}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("strided_slice", "StridedSlice") + .Input(input.node()->name(), 0, DT_INT32) + .Input(begin.node()->name(), 0, DT_INT32) + .Input(end.node()->name(), 0, DT_INT32) + .Input(strides.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Index", DT_INT32) + .Attr("begin_mask", 0x3) + .Attr("end_mask", 0x3) + .Attr("ellipsis_mask", 0) + .Attr("new_axis_mask", 0) + .Attr("shrink_axis_mask", 0x2) + .Finalize(&def)); + absl::Status status; + Node* strided_slice = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, strided_slice, 0); + scope.graph()->AddEdge(begin.node(), 0, strided_slice, 1); + scope.graph()->AddEdge(end.node(), 0, strided_slice, 2); + scope.graph()->AddEdge(strides.node(), 0, strided_slice, 3); + TF_ASSERT_OK(scope.DoShapeInference(strided_slice)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(strided_slice), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 5, 5}, + std::vector{xla::DExpr::Var(43), xla::DExpr::Const(5), + xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "strided_slice_shrink", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(43))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, + StridedSliceNegativeStridePreservesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto begin = ops::Const(scope.WithOpName("begin"), {0, 0}, {2}); + auto end = ops::Const(scope.WithOpName("end"), {0, 0}, {2}); + auto strides = + ops::Const(scope.WithOpName("strides"), {-1, 1}, {2}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("strided_slice", "StridedSlice") + .Input(input.node()->name(), 0, DT_INT32) + .Input(begin.node()->name(), 0, DT_INT32) + .Input(end.node()->name(), 0, DT_INT32) + .Input(strides.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Index", DT_INT32) + .Attr("begin_mask", 0x3) + .Attr("end_mask", 0x3) + .Attr("ellipsis_mask", 0) + .Attr("new_axis_mask", 0) + .Attr("shrink_axis_mask", 0) + .Finalize(&def)); + absl::Status status; + Node* strided_slice = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, strided_slice, 0); + scope.graph()->AddEdge(begin.node(), 0, strided_slice, 1); + scope.graph()->AddEdge(end.node(), 0, strided_slice, 2); + scope.graph()->AddEdge(strides.node(), 0, strided_slice, 3); + TF_ASSERT_OK(scope.DoShapeInference(strided_slice)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(strided_slice), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 5}, + std::vector{xla::DExpr::Var(44), xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "strided_slice_negative", + std::move(graph), args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(44))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); +} + +TEST_F(XlaCompilerDynamicSizesTest, + StridedSliceNegativeStrideTwoScalesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto begin = ops::Const(scope.WithOpName("begin"), {0, 0}, {2}); + auto end = ops::Const(scope.WithOpName("end"), {0, 0}, {2}); + auto strides = + ops::Const(scope.WithOpName("strides"), {-2, 1}, {2}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("strided_slice", "StridedSlice") + .Input(input.node()->name(), 0, DT_INT32) + .Input(begin.node()->name(), 0, DT_INT32) + .Input(end.node()->name(), 0, DT_INT32) + .Input(strides.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Index", DT_INT32) + .Attr("begin_mask", 0x3) + .Attr("end_mask", 0x3) + .Attr("ellipsis_mask", 0) + .Attr("new_axis_mask", 0) + .Attr("shrink_axis_mask", 0) + .Finalize(&def)); + absl::Status status; + Node* strided_slice = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, strided_slice, 0); + scope.graph()->AddEdge(begin.node(), 0, strided_slice, 1); + scope.graph()->AddEdge(end.node(), 0, strided_slice, 2); + scope.graph()->AddEdge(strides.node(), 0, strided_slice, 3); + TF_ASSERT_OK(scope.DoShapeInference(strided_slice)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(strided_slice), 0); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 5}, + std::vector{xla::DExpr::Var(45), xla::DExpr::Const(5)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "strided_slice_negative_two", + std::move(graph), args, &result)); -class DummyReadResourceOp : public XlaOpKernel { - public: - explicit DummyReadResourceOp(OpKernelConstruction* ctx) : XlaOpKernel(ctx) {} - void Compile(XlaOpKernelContext* ctx) override { - ResourceMgr* rm = ctx->op_kernel_context()->resource_manager(); - OP_REQUIRES(ctx, rm, errors::Internal("No resource manager.")); - DummyResourceForTest* dummy; - OP_REQUIRES_OK(ctx, rm->Lookup( - rm->default_container(), "dummy", &dummy)); - dummy->Increment(); - dummy->Unref(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal( + result.outputs[0].shape.get_filled_expression(0), + ((xla::DExpr::Var(45) + xla::DExpr::Const(1)) / xla::DExpr::Const(2)) + .simplify())); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); +} - ctx->SetOutput(0, ctx->Input(0)); - ctx->SetOutput(1, ctx->Input(0)); - } -}; +TEST_F(XlaCompilerDynamicSizesTest, PadAddsToLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto paddings = + ops::Const(scope.WithOpName("paddings"), {1, 2, 0, 0}, {2, 2}); + auto padded = ops::Pad(scope.WithOpName("pad"), input, paddings); + auto retval = ops::_Retval(scope.WithOpName("retval"), padded, 0); -class DummyReadResourceCC { - public: - DummyReadResourceCC(const Scope& scope, const Input& value) { - if (!scope.ok()) return; - auto _value = ops::AsNodeOut(scope, value); - if (!scope.ok()) return; - Node* ret; - const auto unique_name = scope.GetUniqueNameForOp("DummyReadResource"); - auto builder = NodeBuilder(unique_name, "DummyReadResource").Input(_value); - scope.UpdateBuilder(&builder); - scope.UpdateStatus(builder.Finalize(scope.graph(), &ret)); - if (!scope.ok()) return; - scope.UpdateStatus(scope.DoShapeInference(ret)); - if (!scope.ok()) return; - this->output1_ = Output(ret, 0); - this->output2_ = Output(ret, 1); - } + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); - Output output1_; - Output output2_; -}; + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {7, 5}, + std::vector{xla::DExpr::Var(28), xla::DExpr::Const(5)}); -REGISTER_OP("DummyReadResource") - .Input("input: int32") - .Output("output1: int32") - .Output("output2: int32") - .SetShapeFn(shape_inference::UnknownShape) - .Doc(R"doc( -A dummy Op. + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "pad", + std::move(graph), args, &result)); -input: dummy input. -output1: dummy output. -output2: dummy output. -)doc"); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + (xla::DExpr::Var(28) + 3).simplify())); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); +} -REGISTER_XLA_OP(Name("DummyReadResource"), DummyReadResourceOp); +TEST_F(XlaCompilerDynamicSizesTest, SpaceToBatchNDScalesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_FLOAT, 0); + auto block_shape = + ops::Const(scope.WithOpName("block_shape"), {2}, {1}); + auto paddings = + ops::Const(scope.WithOpName("paddings"), {0, 0}, {1, 2}); -// DummyDuplicateOp is present purely to test multiple REGISTER_XLA_OP calls -// on the same Op name below. -class DummyDuplicateOp : public XlaOpKernel { - public: - explicit DummyDuplicateOp(OpKernelConstruction* ctx) : XlaOpKernel(ctx) {} - void Compile(XlaOpKernelContext* ctx) override { - ctx->SetOutput(0, ctx->Input(0)); - } -}; + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("space_to_batch", "SpaceToBatchND") + .Input(input.node()->name(), 0, DT_FLOAT) + .Input(block_shape.node()->name(), 0, DT_INT32) + .Input(paddings.node()->name(), 0, DT_INT32) + .Attr("T", DT_FLOAT) + .Attr("Tblock_shape", DT_INT32) + .Attr("Tpaddings", DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* space_to_batch = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, space_to_batch, 0); + scope.graph()->AddEdge(block_shape.node(), 0, space_to_batch, 1); + scope.graph()->AddEdge(paddings.node(), 0, space_to_batch, 2); + TF_ASSERT_OK(scope.DoShapeInference(space_to_batch)); -REGISTER_OP("DummyDuplicateOp") - .Input("input: int32") - .Output("output: int32") - .Doc(R"doc( -A dummy Op. + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(space_to_batch), 0); -input: dummy input. -output: dummy output. -)doc"); + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); -REGISTER_XLA_OP(Name("DummyDuplicateOp").Device(DEVICE_CPU_XLA_JIT), - DummyDuplicateOp); -REGISTER_XLA_OP(Name("DummyDuplicateOp").Device(DEVICE_GPU_XLA_JIT), - DummyDuplicateOp); + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_FLOAT; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::F32, {4, 8}, + std::vector{xla::DExpr::Var(29), xla::DExpr::Const(8)}); -// Tests compilation and execution of an empty graph. -TEST_F(XlaCompilerTest, EmptyReturnValues) { XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "space_to_batch", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + (xla::DExpr::Const(2) * xla::DExpr::Var(29)) + .simplify())); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); +} + +TEST_F(XlaCompilerDynamicSizesTest, BatchToSpaceNDDividesLeadingExpression) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_FLOAT, 0); + auto block_shape = + ops::Const(scope.WithOpName("block_shape"), {2}, {1}); + auto crops = ops::Const(scope.WithOpName("crops"), {0, 0}, {1, 2}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("batch_to_space", "BatchToSpaceND") + .Input(input.node()->name(), 0, DT_FLOAT) + .Input(block_shape.node()->name(), 0, DT_INT32) + .Input(crops.node()->name(), 0, DT_INT32) + .Attr("T", DT_FLOAT) + .Attr("Tblock_shape", DT_INT32) + .Attr("Tcrops", DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* batch_to_space = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, batch_to_space, 0); + scope.graph()->AddEdge(block_shape.node(), 0, batch_to_space, 1); + scope.graph()->AddEdge(crops.node(), 0, batch_to_space, 2); + TF_ASSERT_OK(scope.DoShapeInference(batch_to_space)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(batch_to_space), 0); std::unique_ptr graph(new Graph(OpRegistry::Global())); - XlaCompiler::CompilationResult result; - TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "add", - std::move(graph), - /*args=*/{}, &result)); + TF_ASSERT_OK(scope.ToGraph(graph.get())); - TF_ASSERT_OK(client_->Execute(*result.computation, {}).status()); + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_FLOAT; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::F32, {8, 4}, + std::vector{(xla::DExpr::Const(2) * xla::DExpr::Var(30)) + .simplify(), + xla::DExpr::Const(4)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "batch_to_space", std::move(graph), args, + &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(30))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(8))); } -// Tests compilation and execution of a graph that adds two tensors. -TEST_F(XlaCompilerTest, Simple) { - // Builds a graph that adds two Tensors. +TEST_F(XlaCompilerDynamicSizesTest, SpaceToDepthScalesDepthExpression) { Scope scope = Scope::NewRootScope().ExitOnError(); - auto a = ops::_Arg(scope.WithOpName("A"), DT_INT32, 0); - auto b = ops::_Arg(scope.WithOpName("B"), DT_INT32, 1); - auto c = ops::Add(scope.WithOpName("C"), a, b); - auto d = ops::_Retval(scope.WithOpName("D"), c, 0); + auto input = ops::_Arg(scope.WithOpName("input"), DT_FLOAT, 0); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("space_to_depth", "SpaceToDepth") + .Input(input.node()->name(), 0, DT_FLOAT) + .Attr("T", DT_FLOAT) + .Attr("block_size", 2) + .Attr("data_format", "NHWC") + .Finalize(&def)); + absl::Status status; + Node* space_to_depth = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, space_to_depth, 0); + TF_ASSERT_OK(scope.DoShapeInference(space_to_depth)); + + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(space_to_depth), 0); + std::unique_ptr graph(new Graph(OpRegistry::Global())); TF_ASSERT_OK(scope.ToGraph(graph.get())); - // Builds a description of the arguments. - std::vector args(2); + std::vector args(1); args[0].kind = XlaCompiler::Argument::kParameter; - args[0].type = DT_INT32; - args[0].shape = TensorShape({2}); - args[1].kind = XlaCompiler::Argument::kParameter; - args[1].type = DT_INT32; - args[1].shape = TensorShape({2}); + args[0].type = DT_FLOAT; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::F32, {5, 8, 8, 3}, + std::vector{xla::DExpr::Const(5), xla::DExpr::Const(8), + xla::DExpr::Const(8), xla::DExpr::Var(35)}); - // Compiles the graph. XlaCompiler compiler(DefaultOptions()); XlaCompiler::CompilationResult result; - TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "add", - std::move(graph), args, &result)); + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "space_to_depth", std::move(graph), args, + &result)); + + xla::DExpr expected_depth = + (xla::DExpr::Const(4) * xla::DExpr::Var(35)).simplify(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Const(5))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(4))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(4))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 3), + expected_depth)); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(0), xla::DExpr::Const(5))); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(1), xla::DExpr::Const(4))); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(2), xla::DExpr::Const(4))); + EXPECT_TRUE(xla::DynExpr::equal(result_shape.expressions(3), expected_depth)); +} - // Tests that the generated computation works. - xla::Literal param0_literal = xla::LiteralUtil::CreateR1({7, 42}); - xla::Literal param1_literal = xla::LiteralUtil::CreateR1({-3, 101}); - std::unique_ptr param0_data = - client_->TransferToServer(param0_literal).value(); - std::unique_ptr param1_data = - client_->TransferToServer(param1_literal).value(); +TEST_F(XlaCompilerDynamicSizesTest, DepthToSpaceScalesSpatialExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto input = ops::_Arg(scope.WithOpName("input"), DT_FLOAT, 0); - std::unique_ptr actual = - client_ - ->Execute(*result.computation, {param0_data.get(), param1_data.get()}) - .value(); - xla::Literal actual_literal = client_->Transfer(*actual).value(); + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("depth_to_space", "DepthToSpace") + .Input(input.node()->name(), 0, DT_FLOAT) + .Attr("T", DT_FLOAT) + .Attr("block_size", 2) + .Attr("data_format", "NHWC") + .Finalize(&def)); + absl::Status status; + Node* depth_to_space = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, depth_to_space, 0); + TF_ASSERT_OK(scope.DoShapeInference(depth_to_space)); - xla::Literal expected0 = xla::LiteralUtil::CreateR1({4, 143}); - xla::Literal expected_literal = xla::LiteralUtil::MakeTuple({&expected0}); - EXPECT_TRUE(xla::LiteralTestUtil::Equal(expected_literal, actual_literal)); -} + auto retval = + ops::_Retval(scope.WithOpName("retval"), Output(depth_to_space), 0); -absl::StatusOr> LoadModuleFromHloProto( - const xla::HloModuleProto& module_proto) { - TF_ASSIGN_OR_RETURN(auto module_config, - xla::HloModule::CreateModuleConfigFromProto( - module_proto, xla::GetDebugOptionsFromFlags())); - return xla::CreateModuleFromProto(module_proto, module_config); + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_FLOAT; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::F32, {5, 4, 4, 12}, + std::vector{xla::DExpr::Const(5), xla::DExpr::Var(37), + xla::DExpr::Const(4), xla::DExpr::Const(12)}); + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "depth_to_space", std::move(graph), args, + &result)); + + xla::DExpr expected_height = + (xla::DExpr::Const(2) * xla::DExpr::Var(37)).simplify(); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Const(5))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + expected_height)); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 2), + xla::DExpr::Const(8))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 3), + xla::DExpr::Const(3))); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(0), xla::DExpr::Const(5))); + EXPECT_TRUE(xla::DynExpr::equal(result_shape.expressions(1), expected_height)); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(2), xla::DExpr::Const(8))); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(3), xla::DExpr::Const(3))); } -// Tests compilation and execution of a graph that adds two tensors with dynamic -// shape parameters. -TEST_F(XlaCompilerTest, SimpleDynamicShapeParameter) { - // Builds a graph that adds two Tensors. +TEST_F(XlaCompilerDynamicSizesTest, RollPreservesExpressions) { Scope scope = Scope::NewRootScope().ExitOnError(); - auto a = ops::_Arg(scope.WithOpName("A"), DT_INT32, 0); - auto b = ops::_Arg(scope.WithOpName("B"), DT_INT32, 1); - auto c = ops::Add(scope.WithOpName("C"), a, b); - auto d = ops::_Retval(scope.WithOpName("D"), c, 0); + auto input = ops::_Arg(scope.WithOpName("input"), DT_INT32, 0); + auto shift = ops::Const(scope.WithOpName("shift"), 2, {}); + auto axis = ops::Const(scope.WithOpName("axis"), 1, {}); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("roll", "Roll") + .Input(input.node()->name(), 0, DT_INT32) + .Input(shift.node()->name(), 0, DT_INT32) + .Input(axis.node()->name(), 0, DT_INT32) + .Attr("T", DT_INT32) + .Attr("Tshift", DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* roll = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + scope.graph()->AddEdge(input.node(), 0, roll, 0); + scope.graph()->AddEdge(shift.node(), 0, roll, 1); + scope.graph()->AddEdge(axis.node(), 0, roll, 2); + TF_ASSERT_OK(scope.DoShapeInference(roll)); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(roll), 0); + std::unique_ptr graph(new Graph(OpRegistry::Global())); TF_ASSERT_OK(scope.ToGraph(graph.get())); - // Builds a description of the arguments. - std::vector args(2); + std::vector args(1); args[0].kind = XlaCompiler::Argument::kParameter; args[0].type = DT_INT32; - args[0].shape = - xla::ShapeUtil::MakeShape(/*element_type=*/xla::S32, /*dimensions=*/{2}, - /*dynamic_dimensions=*/{true}); - args[1].kind = XlaCompiler::Argument::kParameter; - args[1].type = DT_INT32; - args[1].shape = TensorShape(/*dimensions=*/{2}); + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {6, 5}, + std::vector{xla::DExpr::Var(21), xla::DExpr::Const(5)}); - // Compiles the graph. XlaCompiler compiler(DefaultOptions()); XlaCompiler::CompilationResult result; - TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "add", + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), "roll", std::move(graph), args, &result)); - auto hlo = result.computation->proto(); - TF_ASSERT_OK_AND_ASSIGN(auto module, LoadModuleFromHloProto(hlo)); - EXPECT_EQ(module->computation_count(), 1); - EXPECT_TRUE(module->mutable_computation(0) - ->parameter_instruction(0) - ->shape() - .is_dynamic()); + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(21))); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 1), + xla::DExpr::Const(5))); } // Tests compilation of a graph where the _Retval node is not necessarily last @@ -1014,6 +3566,13 @@ FunctionDef FillFn() { {{{"y"}, "Fill", {"dims", "x"}, {{"T", "$T"}}}}); } +FunctionDef IdentityFn() { + return FunctionDefHelper::Define( + "IdentityFn", {"x: T"}, {"y: T"}, + {"T: {float, double, int32, int64}"}, + {{{"y"}, "Identity", {"x"}, {{"T", "$T"}}}}); +} + TEST_F(XlaCompilerTest, FunctionCallWithConstants) { // Certain operations in a function, "Fill" for example, requires the // operator's argument to be a compile-time constant instead of a parameter. @@ -1057,6 +3616,53 @@ TEST_F(XlaCompilerTest, FunctionCallWithConstants) { std::move(graph), args, &result)); } +TEST_F(XlaCompilerDynamicSizesTest, FunctionCallPreservesDynamicExpressions) { + XlaCompiler compiler(DefaultOptions()); + + FunctionDefLibrary flib; + *flib.add_function() = IdentityFn(); + TF_ASSERT_OK(flib_def_->AddFunctionDef(IdentityFn())); + + std::unique_ptr graph(new Graph(OpRegistry::Global())); + Scope scope = Scope::NewRootScope().ExitOnError(); + auto arg = ops::_Arg(scope.WithOpName("arg"), DT_INT32, 0); + TF_EXPECT_OK(scope.graph()->AddFunctionLibrary(flib)); + + NodeDef def; + TF_ASSERT_OK(NodeDefBuilder("identity_fn", "IdentityFn", flib_def_.get()) + .Input(arg.node()->name(), 0, DT_INT32) + .Finalize(&def)); + absl::Status status; + Node* identity_fn = scope.graph()->AddNode(def, &status); + TF_ASSERT_OK(status); + TF_ASSERT_OK(scope.DoShapeInference(identity_fn)); + scope.graph()->AddEdge(arg.node(), 0, identity_fn, 0); + + auto retval = ops::_Retval(scope.WithOpName("retval"), Output(identity_fn), 0); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(1); + args[0].kind = XlaCompiler::Argument::kParameter; + args[0].type = DT_INT32; + args[0].shape = xla::ShapeUtil::MakeShape( + xla::S32, {9}, std::vector{xla::DExpr::Var(2)}); + + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "identity_function", std::move(graph), + args, &result)); + + ASSERT_EQ(result.outputs.size(), 1); + EXPECT_TRUE(xla::DynExpr::equal(result.outputs[0].shape.get_filled_expression( + 0), + xla::DExpr::Var(2))); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_TRUE( + xla::DynExpr::equal(result_shape.expressions(0), xla::DExpr::Var(2))); +} + // Tests CompileFunction with a local function lookup failing, fails with // informative error about both lookups. TEST_F(XlaCompilerTest, LocalFunctionWithWrongArgumentsFail) { diff --git a/tensorflow/core/framework/tensor_shape.cc b/tensorflow/core/framework/tensor_shape.cc index fd2606224b3429..0ceeb8c9b169f6 100644 --- a/tensorflow/core/framework/tensor_shape.cc +++ b/tensorflow/core/framework/tensor_shape.cc @@ -30,8 +30,6 @@ namespace tensorflow { namespace { -const bool kTensorShapeExpressionsEnabled = TensorShapeExpressionsEnabled(); - xla::DExpr DExprFromProto(const ExpressionProto& proto) { switch (proto.node_type_case()) { case ExpressionProto::kConstantValue: @@ -246,7 +244,7 @@ TensorShapeBase::TensorShapeBase(const TensorShapeProto& proto) { for (const auto& d : proto.dim()) { AddDim(d.size()); } - if (kTensorShapeExpressionsEnabled) { + if (TensorShapeExpressionsEnabled()) { for (const auto& e : proto.expressions()) { AddExpression(DExprFromProto(e)); } @@ -290,7 +288,7 @@ absl::Status TensorShapeBase::BuildTensorShapeBase( } } } - if (kTensorShapeExpressionsEnabled) { + if (TensorShapeExpressionsEnabled()) { for (const auto& e : proto.expressions()) { out->AddExpression(DExprFromProto(e)); } @@ -480,7 +478,7 @@ void TensorShapeRep::Clear() { } void TensorShapeRep::set_expression(int d, xla::DExpr expr) { - if (!kTensorShapeExpressionsEnabled) { + if (!TensorShapeExpressionsEnabled()) { expressions_.clear(); return; } @@ -493,7 +491,7 @@ void TensorShapeRep::set_expression(int d, xla::DExpr expr) { } void TensorShapeRep::AddExpression(xla::DExpr expr) { - if (!kTensorShapeExpressionsEnabled) { + if (!TensorShapeExpressionsEnabled()) { return; } CHECK_LT(expressions_.size(), ndims_byte()); @@ -503,7 +501,7 @@ void TensorShapeRep::AddExpression(xla::DExpr expr) { } void TensorShapeRep::set_expressions(std::vector exprs) { - if (!kTensorShapeExpressionsEnabled) { + if (!TensorShapeExpressionsEnabled()) { expressions_.clear(); return; } @@ -926,7 +924,7 @@ void TensorShapeBase::AsProto(TensorShapeProto* proto) const { for (int i = 0; i < dims(); i++) { proto->add_dim()->set_size(dim_size(i)); } - if (kTensorShapeExpressionsEnabled) { + if (TensorShapeExpressionsEnabled()) { for (int i = 0; i < get_expressions().size(); i++) { ExpressionProto* eproto = proto->add_expressions(); ExprToProto(get_expression(i), eproto); @@ -993,7 +991,7 @@ string TensorShapeRep::DebugString(const TensorShapeProto& proto) { first = false; } strings::StrAppend(&s, "]"); - if (kTensorShapeExpressionsEnabled) { + if (TensorShapeExpressionsEnabled()) { strings::StrAppend(&s, "<"); first = true; for (const auto& e : proto.expressions()) { diff --git a/tensorflow/core/framework/tensor_shape_expr.cc b/tensorflow/core/framework/tensor_shape_expr.cc index baca4df733d63f..686992d4c9881a 100644 --- a/tensorflow/core/framework/tensor_shape_expr.cc +++ b/tensorflow/core/framework/tensor_shape_expr.cc @@ -19,13 +19,25 @@ bool ParseTensorShapeExpressionsEnabled() { return tf_xla_enable_dynamic_sizes; } +std::optional& TensorShapeExpressionsEnabledOverride() { + static auto* enabled_override = new std::optional(); + return *enabled_override; +} + } // namespace bool TensorShapeExpressionsEnabled() { + if (TensorShapeExpressionsEnabledOverride().has_value()) { + return *TensorShapeExpressionsEnabledOverride(); + } static const bool enabled = ParseTensorShapeExpressionsEnabled(); return enabled; } +void SetTensorShapeExpressionsEnabledForTesting(std::optional enabled) { + TensorShapeExpressionsEnabledOverride() = enabled; +} + bool IsDynamicDimExpr(const ExpressionProto& proto) { switch (proto.node_type_case()) { case ExpressionProto::kVariableId: diff --git a/tensorflow/core/framework/tensor_shape_expr.h b/tensorflow/core/framework/tensor_shape_expr.h index 1c215fda268dcf..0a2497c03b808e 100644 --- a/tensorflow/core/framework/tensor_shape_expr.h +++ b/tensorflow/core/framework/tensor_shape_expr.h @@ -3,6 +3,7 @@ #include #include +#include #include #include @@ -218,6 +219,10 @@ DimExpr* SimplifyExpr(DimExpr* expr, // Shape-expression support follows the `tf_xla_enable_dynamic_sizes` flag. bool TensorShapeExpressionsEnabled(); +// Overrides TensorShapeExpressionsEnabled for tests. Passing std::nullopt +// restores the default environment-derived behavior. +void SetTensorShapeExpressionsEnabledForTesting(std::optional enabled); + // Returns true if the expression proto depends on a symbolic variable. bool IsDynamicDimExpr(const ExpressionProto& proto); diff --git a/third_party/xla/xla/service/shape_inference.cc b/third_party/xla/xla/service/shape_inference.cc index 4226e901326321..565093a7d81a62 100644 --- a/third_party/xla/xla/service/shape_inference.cc +++ b/third_party/xla/xla/service/shape_inference.cc @@ -77,6 +77,15 @@ bool CompatibleDimensionSizes(int64_t size_a, int64_t size_b) { size_a == size_b; } +DExpr SymbolicElementsIn(const Shape& shape) { + DExpr product = DExpr::Const(1); + for (int64_t i = 0; i < shape.dimensions_size(); ++i) { + const DExpr& expr = shape.expressions(i); + product = product * (expr ? expr : DExpr::Const(shape.dimensions(i))); + } + return product.simplify(); +} + absl::Status ExpectArray(const Shape& shape, absl::string_view op_type) { if (!shape.IsArray()) { return InvalidArgument("Expected array argument for %s, but got %s.", @@ -217,11 +226,14 @@ absl::StatusOr InferWindowOutputShape(const Shape& base_shape, window.DebugString()); } - if (IsUnboundedDynamicSize(ShapeUtil::GetDimension(base_shape, i))) { + const int64_t input_dimension = ShapeUtil::GetDimension(base_shape, i); + const DExpr& input_expression = base_shape.expressions(i); + + if (IsUnboundedDynamicSize(input_dimension)) { output_dimensions[i] = Shape::kUnboundedSize; } else { const int64_t dilated_base = window_util::DilatedBound( - ShapeUtil::GetDimension(base_shape, i), dim.base_dilation()); + input_dimension, dim.base_dilation()); const int64_t padded_dilated_base = dim.padding_low() + dilated_base + dim.padding_high(); const int64_t dilated_window = @@ -231,7 +243,20 @@ absl::StatusOr InferWindowOutputShape(const Shape& base_shape, padded_dilated_base, dilated_window, dim.stride()); } output_is_dynamic[i] = base_shape.is_dynamic_dimension(i); - output_expressions[i] = base_shape.expressions(i); + if (input_expression && input_expression->is_constant()) { + output_expressions[i] = DExpr::Const(output_dimensions[i]); + continue; + } + + DExpr dilated_base_expr = + ((dim.base_dilation() * (input_expression - 1)) + 1).simplify(); + DExpr padded_dilated_base_expr = + (dilated_base_expr + dim.padding_low() + dim.padding_high()).simplify(); + DExpr dilated_window_expr = + DExpr::Const(window_util::DilatedBound(dim.size(), dim.window_dilation())); + output_expressions[i] = + (((padded_dilated_base_expr - dilated_window_expr) / dim.stride()) + 1) + .simplify(); } return ShapeUtil::MakeValidatedShape(element_type, output_dimensions, @@ -2460,11 +2485,13 @@ ShapeInference::InferScalarBroadcastShape(absl::Span shapes) { std::vector dynamic_dimensions(input_spatial_dims.size()); std::vector expressions(input_spatial_dims.size()); - for (auto it = input_spatial_dims.begin(); it != input_spatial_dims.end(); - ++it) { - dynamic_dimensions[it - input_spatial_dims.begin()] = - IsUnboundedDynamicSize(*it); - expressions[it - input_spatial_dims.begin()] = DExpr::Unknown(70); + for (int i = 0; i < input_spatial_dims.size(); ++i) { + const int64_t input_spatial_dimension = + dnums.input_spatial_dimensions(i); + dynamic_dimensions[i] = IsUnboundedDynamicSize(input_spatial_dims[i]); + expressions[i] = lhs.expressions(input_spatial_dimension) + ? lhs.expressions(input_spatial_dimension) + : DExpr::Const(input_spatial_dims[i]); } Shape base_shape = ShapeUtil::MakeShape( lhs.element_type(), input_spatial_dims, dynamic_dimensions, @@ -3900,6 +3927,16 @@ ShapeInference::InferCollectivePermuteDoneShape(const Shape& operand_shape) { ShapeUtil::ElementsIn(inferred_shape), ShapeUtil::HumanString(inferred_shape)); } + if (!expressions.empty()) { + DExpr input_elements = SymbolicElementsIn(operand); + DExpr output_elements = SymbolicElementsIn(inferred_shape); + if (!DynExpr::equal(input_elements.get(), output_elements.get())) { + return InvalidArgument( + "Reshape operation has mismatched symbolic element counts: " + "from=%s to=%s.", + ShapeUtil::HumanString(operand), ShapeUtil::HumanString(inferred_shape)); + } + } std::vector indices(operand.dimensions_size()); std::iota(indices.begin(), indices.end(), 0); @@ -4127,6 +4164,15 @@ ShapeInference::InferCollectivePermuteDoneShape(const Shape& operand_shape) { on_true.is_dynamic_dimension(dimension) || on_false.is_dynamic_dimension(dimension)); } + if (!DynExpr::equal(on_true.expressions(dimension), + on_false.expressions(dimension))) { + return InvalidArgument( + "Select operands have mismatched expressions in dimension %d: " + "on_true=%s, on_false=%s.", + dimension, ShapeUtil::HumanString(on_true), + ShapeUtil::HumanString(on_false)); + } + result.set_expression(dimension, on_true.expressions(dimension)); } if (result.has_layout()) { result.mutable_layout()->set_element_size_in_bits( diff --git a/third_party/xla/xla/service/shape_inference_test.cc b/third_party/xla/xla/service/shape_inference_test.cc index 39b71da5f23ca8..7efbfbb4684f5a 100644 --- a/third_party/xla/xla/service/shape_inference_test.cc +++ b/third_party/xla/xla/service/shape_inference_test.cc @@ -205,6 +205,34 @@ TEST_F(ShapeInferenceTest, SelectArrayPredBetweenArrays) { ASSERT_TRUE(ShapeUtil::Equal(matrix_64_48_, *inferred_shape)); } +TEST_F(ShapeInferenceTest, SelectPreservesExpressionsFromOperands) { + const Shape pred = ShapeUtil::MakeShape(PRED, {8, 5}); + const Shape on_true = ShapeUtil::MakeShape( + F32, {8, 5}, std::vector{DExpr::Var(33), DExpr::Const(5)}); + const Shape on_false = ShapeUtil::MakeShape( + F32, {8, 5}, std::vector{DExpr::Var(33), DExpr::Const(5)}); + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferTernaryOpShape(HloOpcode::kSelect, pred, on_true, + on_false)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(33))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(5))); +} + +TEST_F(ShapeInferenceTest, SelectRejectsMismatchedOperandExpressions) { + const Shape pred = ShapeUtil::MakeShape(PRED, {}); + const Shape on_true = ShapeUtil::MakeShape( + F32, {8, 5}, std::vector{DExpr::Var(33), DExpr::Const(5)}); + const Shape on_false = ShapeUtil::MakeShape( + F32, {8, 5}, std::vector{DExpr::Var(34), DExpr::Const(5)}); + const absl::StatusOr inferred_shape = + ShapeInference::InferTernaryOpShape(HloOpcode::kSelect, pred, on_true, + on_false); + ASSERT_FALSE(inferred_shape.ok()); + EXPECT_THAT(inferred_shape.status().message(), + HasSubstr("mismatched expressions in dimension 0")); +} + TEST_F(ShapeInferenceTest, SelectBadShapes) { const absl::StatusOr inferred_shape_error1 = ShapeInference::InferTernaryOpShape(HloOpcode::kSelect, pred_, @@ -448,6 +476,63 @@ TEST_F(ShapeInferenceTest, ReduceWindowInHalf) { ShapeUtil::Equal(ShapeUtil::MakeShape(F32, {4, 4}), *inferred_shape)); } +TEST_F(ShapeInferenceTest, ReduceWindowBuildsWindowedExpressions) { + const Shape matrix_shape = ShapeUtil::MakeShape( + F32, {8, 8}, std::vector{DExpr::Var(1), DExpr::Var(2)}); + Window window; + WindowDimension dim; + dim.set_size(2); + dim.set_stride(2); + dim.set_padding_low(0); + dim.set_padding_high(0); + dim.set_window_dilation(1); + dim.set_base_dilation(1); + *window.add_dimensions() = dim; + *window.add_dimensions() = dim; + const Shape init_value_shape = ShapeUtil::MakeShape(F32, {}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReduceWindowShape(matrix_shape, init_value_shape, + window)); + + EXPECT_TRUE(ShapeUtil::Equal( + inferred_shape, ShapeUtil::MakeShape( + F32, {4, 4}, + std::vector{((DExpr::Var(1) - 2) / 2) + 1, + ((DExpr::Var(2) - 2) / 2) + 1}))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), + ((DExpr::Var(1) - 2) / 2) + 1)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), + ((DExpr::Var(2) - 2) / 2) + 1)); +} + +TEST_F(ShapeInferenceTest, ReduceWindowBuildsPaddedStridedExpressions) { + const Shape vector_shape = + ShapeUtil::MakeShape(F32, {11}, std::vector{DExpr::Var(3)}); + Window window; + WindowDimension dim; + dim.set_size(3); + dim.set_stride(2); + dim.set_padding_low(1); + dim.set_padding_high(1); + dim.set_window_dilation(1); + dim.set_base_dilation(1); + *window.add_dimensions() = dim; + const Shape init_value_shape = ShapeUtil::MakeShape(F32, {}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReduceWindowShape(vector_shape, init_value_shape, + window)); + + const DExpr expected = (((DExpr::Var(3) - 1) / 2) + 1).simplify(); + EXPECT_TRUE(ShapeUtil::Equal( + inferred_shape, + ShapeUtil::MakeShape(F32, {6}, std::vector{expected}))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), expected)); +} + TEST_F(SelectAndScatterShapeInferenceTest, SelectAndScatterProperShapes) { const absl::StatusOr inferred_shape_ok = ShapeInference::InferSelectAndScatterShape( @@ -721,6 +806,58 @@ TEST_F(ShapeInferenceTest, ConvolveWithBaseDilation) { *inferred_shape)); } +TEST_F(ShapeInferenceTest, ConvolveBuildsBatchAndSpatialExpressions) { + ConvolutionDimensionNumbers dnums; + const Shape lhs_shape = ShapeUtil::MakeShape( + F32, {5, 11, 13, 3}, + std::vector{DExpr::Var(4), DExpr::Var(5), DExpr::Var(6), + DExpr::Const(3)}); + dnums.set_input_batch_dimension(0); + dnums.set_output_batch_dimension(0); + dnums.add_input_spatial_dimensions(1); + dnums.add_output_spatial_dimensions(1); + dnums.add_input_spatial_dimensions(2); + dnums.add_output_spatial_dimensions(2); + dnums.set_input_feature_dimension(3); + dnums.set_output_feature_dimension(3); + + const Shape rhs_shape = ShapeUtil::MakeShape(F32, {3, 3, 3, 7}); + dnums.add_kernel_spatial_dimensions(0); + dnums.add_kernel_spatial_dimensions(1); + dnums.set_kernel_input_feature_dimension(2); + dnums.set_kernel_output_feature_dimension(3); + + Window window; + auto* dim0 = window.add_dimensions(); + dim0->set_size(3); + dim0->set_stride(2); + dim0->set_padding_low(1); + dim0->set_padding_high(1); + dim0->set_window_dilation(1); + dim0->set_base_dilation(1); + auto* dim1 = window.add_dimensions(); + *dim1 = *dim0; + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferConvolveShape( + lhs_shape, rhs_shape, /*feature_group_count=*/1, + /*batch_group_count=*/1, window, dnums, + /*preferred_element_type=*/std::nullopt)); + + const DExpr expected_h = (((DExpr::Var(5) - 1) / 2) + 1).simplify(); + const DExpr expected_w = (((DExpr::Var(6) - 1) / 2) + 1).simplify(); + EXPECT_TRUE(ShapeUtil::Equal( + inferred_shape, + ShapeUtil::MakeShape(F32, {5, 6, 7, 7}, + std::vector{DExpr::Var(4), expected_h, + expected_w, DExpr::Const(7)}))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(4))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), expected_h)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(2), expected_w)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(3), DExpr::Const(7))); +} + TEST_F(ShapeInferenceTest, ConvolveDimensionNumbersOverlapError) { // Dimension order for this test: batch, feature, x0, x1 const Shape lhs_shape = ShapeUtil::MakeShape(F32, {10, 11, 3, 4}); @@ -1285,6 +1422,18 @@ TEST_F(ShapeInferenceTest, MapWithDifferentInputTypes) { EXPECT_TRUE(ShapeUtil::Equal(expected, *inferred_shape)); } +TEST_F(ShapeInferenceTest, MapPreservesExpressions) { + const Shape arg = ShapeUtil::MakeShape( + F32, {20, 7}, std::vector{true, false}, + std::vector{DExpr::Var(11), DExpr::Const(7)}); + ProgramShape to_apply = ShapeUtil::MakeProgramShape({f32_}, f32_); + TF_ASSERT_OK_AND_ASSIGN(const Shape inferred_shape, + ShapeInference::InferMapShape({&arg}, to_apply, + {0, 1})); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(11))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(7))); +} + TEST_F(ReduceShapeInferenceTest, ReduceVectorToScalar) { ExpectInferredReduceShape(f32_, ShapeUtil::MakeShape(F32, {128}), /*dimensions_to_reduce=*/{0}); @@ -1532,6 +1681,23 @@ TEST_F(ShapeInferenceTest, InferSliceWithDynamicDimensions) { *inferred_shape)); } +TEST_F(ShapeInferenceTest, InferSliceBuildsExpressionFromSymbolicBounds) { + const Shape vector_shape = + ShapeUtil::MakeShape(F32, {16}, std::vector{DExpr::Var(1)}); + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferSliceShape( + vector_shape, /*starts=*/{3}, /*limits=*/{8}, /*strides=*/{1}, + /*start_exprs=*/{DExpr::Const(3)}, + /*limit_exprs=*/{DExpr::Var(2)})); + + EXPECT_TRUE(ShapeUtil::Equal( + inferred_shape, ShapeUtil::MakeShape( + F32, {5}, std::vector{DExpr::Var(2) - 3}))); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(2) - 3)); +} + TEST_F(ShapeInferenceTest, InferSliceShapeRank2WithStrides) { const Shape matrix_shape = ShapeUtil::MakeShape(F32, {128, 64}); const absl::StatusOr inferred_shape = @@ -1587,6 +1753,26 @@ TEST_F(ShapeInferenceTest, InferConstIndexShape) { ASSERT_TRUE(ShapeUtil::Equal(s32_, *inferred1_status)); } +TEST_F(ShapeInferenceTest, InferConstIndexPreservesExpressions) { + const Shape lhs = ShapeUtil::MakeShape( + F32, {8, 5}, std::vector{DExpr::Var(24), DExpr::Const(5)}); + const Shape rhs = ShapeUtil::MakeShape( + S32, {3, 7}, std::vector{DExpr::Const(3), DExpr::Var(25)}); + const Shape tuple_shape = ShapeUtil::MakeTupleShape({lhs, rhs}); + + TF_ASSERT_OK_AND_ASSIGN(const Shape inferred0, + ShapeInference::InferGetTupleElementShape( + tuple_shape, /*index=*/0)); + TF_ASSERT_OK_AND_ASSIGN(const Shape inferred1, + ShapeInference::InferGetTupleElementShape( + tuple_shape, /*index=*/1)); + + EXPECT_TRUE(DynExpr::equal(inferred0.expressions(0), DExpr::Var(24))); + EXPECT_TRUE(DynExpr::equal(inferred0.expressions(1), DExpr::Const(5))); + EXPECT_TRUE(DynExpr::equal(inferred1.expressions(0), DExpr::Const(3))); + EXPECT_TRUE(DynExpr::equal(inferred1.expressions(1), DExpr::Var(25))); +} + TEST_F(ShapeInferenceTest, InferTupleElementShapeOutOfBound) { const Shape tuple_shape = ShapeUtil::MakeTupleShape({f32_, s32_}); const absl::StatusOr inferredNegative_status = @@ -1677,6 +1863,153 @@ TEST_F(ShapeInferenceTest, UnchangedDimension) { *status); } +TEST_F(ShapeInferenceTest, ReshapePreservesProvidedExpressions) { + const Shape operand = ShapeUtil::MakeShape( + F32, {6, 10}, std::vector{DExpr::Const(6), DExpr::Var(1)}); + const Shape expected = ShapeUtil::MakeShape( + F32, {2, 3, 10}, + std::vector{DExpr::Const(2), DExpr::Const(3), DExpr::Var(1)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReshapeShape( + operand, expected.dimensions(), + /*inferred_dimension=*/-1, expected.expressions())); + + EXPECT_TRUE(ShapeUtil::Equal(inferred_shape, expected)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Const(2))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(3))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(2), DExpr::Var(1))); +} + +TEST_F(ShapeInferenceTest, ReshapeCombinesLeadingSymbolicWithStaticFactor) { + const Shape operand = ShapeUtil::MakeShape( + F32, {5, 16, 32}, + std::vector{DExpr::Var(1), DExpr::Const(16), DExpr::Const(32)}); + const Shape expected = ShapeUtil::MakeShape( + F32, {80, 32}, + std::vector{16 * DExpr::Var(1), DExpr::Const(32)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReshapeShape( + operand, expected.dimensions(), + /*inferred_dimension=*/-1, expected.expressions())); + + EXPECT_TRUE(ShapeUtil::Equal(inferred_shape, expected)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), + 16 * DExpr::Var(1))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(32))); +} + +TEST_F(ShapeInferenceTest, ReshapeCollapsesTwoStaticDimsIntoSymbolicExtent) { + const Shape operand = ShapeUtil::MakeShape( + F32, {5, 4, 8}, + std::vector{DExpr::Var(1), DExpr::Const(4), DExpr::Const(8)}); + const Shape expected = + ShapeUtil::MakeShape(F32, {160}, std::vector{32 * DExpr::Var(1)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReshapeShape( + operand, expected.dimensions(), + /*inferred_dimension=*/-1, expected.expressions())); + + EXPECT_TRUE(ShapeUtil::Equal(inferred_shape, expected)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), + 32 * DExpr::Var(1))); +} + +TEST_F(ShapeInferenceTest, ReshapeSplitsSymbolicExtentByStaticFactor) { + const Shape operand = ShapeUtil::MakeShape( + F32, {80, 8}, std::vector{DExpr::Var(1), DExpr::Const(8)}); + const Shape expected = ShapeUtil::MakeShape( + F32, {40, 16}, + std::vector{DExpr::Var(1) / 2, DExpr::Const(16)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReshapeShape( + operand, expected.dimensions(), + /*inferred_dimension=*/-1, expected.expressions())); + + EXPECT_TRUE(ShapeUtil::Equal(inferred_shape, expected)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), + DExpr::Var(1) / 2)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(16))); +} + +TEST_F(ShapeInferenceTest, ReshapeSplitsAndCollapsesSymbolicExtent) { + const Shape operand = ShapeUtil::MakeShape( + F32, {20, 8, 4}, + std::vector{DExpr::Var(1), DExpr::Const(8), DExpr::Const(4)}); + const Shape expected = ShapeUtil::MakeShape( + F32, {10, 64}, + std::vector{DExpr::Var(1) / 2, DExpr::Const(64)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReshapeShape( + operand, expected.dimensions(), + /*inferred_dimension=*/-1, expected.expressions())); + + EXPECT_TRUE(ShapeUtil::Equal(inferred_shape, expected)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), + DExpr::Var(1) / 2)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(64))); +} + +TEST_F(ShapeInferenceTest, ReshapeRejectsIncorrectCollapsedExpression) { + const Shape operand = ShapeUtil::MakeShape( + F32, {5, 16, 32}, + std::vector{DExpr::Var(1), DExpr::Const(16), DExpr::Const(32)}); + const Shape incorrect = ShapeUtil::MakeShape( + F32, {80, 32}, + std::vector{8 * DExpr::Var(1), DExpr::Const(32)}); + + const absl::StatusOr status = ShapeInference::InferReshapeShape( + operand, incorrect.dimensions(), + /*inferred_dimension=*/-1, incorrect.expressions()); + + ASSERT_FALSE(status.ok()); + EXPECT_THAT(status.status().message(), + HasSubstr("Reshape operation has mismatched symbolic element " + "counts")); +} + +TEST_F(ShapeInferenceTest, + ReshapeRejectsIncorrectSplitAndCollapseExpression) { + const Shape operand = ShapeUtil::MakeShape( + F32, {20, 8, 4}, + std::vector{DExpr::Var(1), DExpr::Const(8), DExpr::Const(4)}); + const Shape incorrect = ShapeUtil::MakeShape( + F32, {10, 64}, + std::vector{DExpr::Var(1), DExpr::Const(64)}); + + const absl::StatusOr status = ShapeInference::InferReshapeShape( + operand, incorrect.dimensions(), + /*inferred_dimension=*/-1, incorrect.expressions()); + + ASSERT_FALSE(status.ok()); + EXPECT_THAT(status.status().message(), + HasSubstr("Reshape operation has mismatched symbolic element " + "counts")); +} + +TEST_F(ShapeInferenceTest, ReshapeWithSymbolicOperandRequiresExpressions) { + const Shape operand = ShapeUtil::MakeShape( + F32, {10, 6}, std::vector{DExpr::Var(1), DExpr::Const(6)}); + + const absl::StatusOr status = + ShapeInference::InferReshapeShape(operand, {2, 5, 6}, + /*inferred_dimension=*/-1, + /*expressions=*/{}); + + ASSERT_FALSE(status.ok()); + EXPECT_THAT(status.status().message(), + HasSubstr("Expressions is empty but operand is dynamic")); +} + TEST_F(ShapeInferenceTest, InferDynamicBroadcast) { // CHECK: // %broadcast = s32[15,<=15]{1,0} broadcast(s32[<=15]{0}), dimensions={1} @@ -1689,6 +2022,23 @@ TEST_F(ShapeInferenceTest, InferDynamicBroadcast) { *inferred_shape); } +TEST_F(ShapeInferenceTest, BroadcastInDimPreservesMappedExpressions) { + const Shape operand = ShapeUtil::MakeShape( + F32, {2, 4}, std::vector{true, true}, + {DExpr::Var(1), DExpr::Var(2)}); + const Shape output = ShapeUtil::MakeShape( + F32, {2, 3, 4}, std::vector{true, false, true}, + {DExpr::Var(1), DExpr::Const(3), DExpr::Var(2)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferBroadcastShape(operand, output, + /*broadcast_dimensions=*/{0, 2})); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(1))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(3))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(2), DExpr::Var(2))); +} + TEST_F(ShapeInferenceTest, BroadcastScalar) { for (auto element_type : {F32, U32, S8}) { const Shape scalar_shape = ShapeUtil::MakeShape(element_type, {}); @@ -2184,6 +2534,27 @@ TEST_F(ShapeInferenceTest, SparseDotMetadata) { ShapeUtil::Equal(inferred_shape, ShapeUtil::MakeShape(U16, {5, 10, 2}))); } +TEST_F(ShapeInferenceTest, SparseDotMetadataPreservesNonSparseExpressions) { + DotDimensionNumbers dot_dnums; + dot_dnums.add_lhs_batch_dimensions(0); + dot_dnums.add_lhs_contracting_dimensions(2); + SparsityDescriptor sparsity_descriptor; + sparsity_descriptor.set_type(SparsityType::SPARSITY_STRUCTURED_N_M); + sparsity_descriptor.set_n(2); + sparsity_descriptor.set_m(4); + sparsity_descriptor.set_index(0); + sparsity_descriptor.set_dimension(2); + + const Shape operand = ShapeUtil::MakeShape( + F32, {5, 10, 16}, + std::vector{DExpr::Var(16), DExpr::Const(10), DExpr::Const(16)}); + TF_ASSERT_OK_AND_ASSIGN(const Shape inferred_shape, + ShapeInference::InferSparseDotMetadataShape( + operand, dot_dnums, sparsity_descriptor)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(16))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(10))); +} + // mode 1 : [m,k], [g,k,n], [g] -> [m,n] TEST_F(ShapeInferenceTest, RaggedDotRaggedNonContracting) { const Shape lhs_shape = ShapeUtil::MakeShape(F32, {11, 5}); @@ -2233,6 +2604,31 @@ TEST_F(ShapeInferenceTest, RaggedDotRaggedContracting) { << " expected: " << ShapeUtil::HumanString(output_shape); } +TEST_F(ShapeInferenceTest, RaggedDotPreservesGroupExpression) { + const Shape lhs_shape = ShapeUtil::MakeShape( + F32, {11, 5}, std::vector{DExpr::Const(11), DExpr::Const(5)}); + const Shape rhs_shape = ShapeUtil::MakeShape( + F32, {5, 7}, std::vector{DExpr::Const(5), DExpr::Const(7)}); + const Shape group_sizes_shape = + ShapeUtil::MakeShape(U32, {3}, std::vector{DExpr::Var(17)}); + + DotDimensionNumbers dot_dnums; + dot_dnums.add_lhs_contracting_dimensions(1); + dot_dnums.add_rhs_contracting_dimensions(0); + RaggedDotDimensionNumbers ragged_dot_dnums; + *ragged_dot_dnums.mutable_dot_dimension_numbers() = dot_dnums; + ragged_dot_dnums.add_lhs_ragged_dimensions(1); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferRaggedDotOpShape( + lhs_shape, rhs_shape, group_sizes_shape, ragged_dot_dnums, + /*preferred_element_type=*/std::nullopt)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(17))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(11))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(2), DExpr::Const(7))); +} + // mode 3 : [b,m,k], [b,k,n], [g] -> [b,m,n] TEST_F(ShapeInferenceTest, RaggedDotRaggedBatch) { const Shape lhs_shape = ShapeUtil::MakeShape(F32, {3, 11, 5}); @@ -2739,6 +3135,54 @@ TEST_F(ShapeInferenceTest, BinOpBroadcastMatrixVector) { ASSERT_FALSE(inferred_shape_mismatch.ok()); } +TEST_F(ShapeInferenceTest, InDimBroadcastPreservesMappedExpressions) { + const Shape smaller = ShapeUtil::MakeShape( + F32, {2, 3}, + std::vector{DExpr::Var(18), DExpr::Const(3)}); + const Shape larger = ShapeUtil::MakeShape( + F32, {2, 4, 3}, + std::vector{DExpr::Const(2), DExpr::Const(4), DExpr::Const(3)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferBinaryOpShape(HloOpcode::kAdd, larger, smaller, + {0, 2})); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(18))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(4))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(2), DExpr::Const(3))); +} + +TEST_F(ShapeInferenceTest, DegenerateBroadcastUsesNonUnitExpressions) { + const Shape lhs = ShapeUtil::MakeShape( + F32, {1, 3}, + std::vector{DExpr::Const(1), DExpr::Const(3)}); + const Shape rhs = ShapeUtil::MakeShape( + F32, {5, 3}, + std::vector{DExpr::Var(19), DExpr::Const(3)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferBinaryOpShape(HloOpcode::kAdd, lhs, rhs, {})); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(19))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(3))); +} + +TEST_F(ShapeInferenceTest, ElementwiseBinaryBroadcastPreservesExpressions) { + const Shape lhs = ShapeUtil::MakeShape( + F32, {2, 4, 3}, + std::vector{DExpr::Const(2), DExpr::Const(4), DExpr::Const(3)}); + const Shape rhs = ShapeUtil::MakeShape( + F32, {2, 3}, + std::vector{DExpr::Var(20), DExpr::Const(3)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferBinaryOpShape(HloOpcode::kAdd, lhs, rhs, {0, 2})); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(20))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(4))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(2), DExpr::Const(3))); +} + TEST_F(ShapeInferenceTest, BinOpBroadcastCubeMatrix) { // Test variations of broadcasting a matrix for a binary add with a cube. const Shape cube = ShapeUtil::MakeShape(F32, {16, 8, 4}); @@ -2910,6 +3354,26 @@ TEST_F(ShapeInferenceTest, ConcatenateWithDynamicShapes) { *inferred_shape)); } +TEST_F(ShapeInferenceTest, ConcatenateAddsConcatDimensionExpressions) { + const Shape lhs = ShapeUtil::MakeShape( + F32, {2, 5}, std::vector{DExpr::Var(1), DExpr::Const(5)}); + const Shape rhs = ShapeUtil::MakeShape( + F32, {3, 5}, std::vector{DExpr::Var(2), DExpr::Const(5)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferConcatOpShape({&lhs, &rhs}, /*dimension=*/0)); + + EXPECT_TRUE(ShapeUtil::Equal( + inferred_shape, + ShapeUtil::MakeShape(F32, {5, 5}, + std::vector{DExpr::Var(1) + DExpr::Var(2), + DExpr::Const(5)}))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), + DExpr::Var(1) + DExpr::Var(2))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(5))); +} + // Tests for the concatenate instruction with proper shapes. TEST_F(ShapeInferenceTest, ConcatenateWithCorrectShapes) { const absl::StatusOr inferred_shape_1 = @@ -3011,6 +3475,64 @@ TEST_F(ShapeInferenceTest, Pad) { HasSubstr("negative size for dimension 1")); } +TEST_F(ShapeInferenceTest, PadAddsConstantOffsetToExpressions) { + const Shape input_shape = ShapeUtil::MakeShape( + F32, {4, 5}, std::vector{DExpr::Var(1), DExpr::Var(2)}); + const Shape padding_value_shape = ShapeUtil::MakeShape(F32, {}); + PaddingConfig padding_config; + auto* dimension0 = padding_config.add_dimensions(); + dimension0->set_edge_padding_low(1); + dimension0->set_edge_padding_high(2); + dimension0->set_interior_padding(1); + auto* dimension1 = padding_config.add_dimensions(); + dimension1->set_edge_padding_low(0); + dimension1->set_edge_padding_high(4); + dimension1->set_interior_padding(0); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferPadShape(input_shape, padding_value_shape, + padding_config)); + + EXPECT_TRUE(ShapeUtil::Equal( + inferred_shape, ShapeUtil::MakeShape( + F32, {10, 9}, + std::vector{DExpr::Var(1) + 6, + DExpr::Var(2) + 4}))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(1) + 6)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Var(2) + 4)); +} + +TEST_F(ShapeInferenceTest, PadBuildsExpressionsForTwoSymbolicDimensions) { + const Shape input_shape = ShapeUtil::MakeShape( + F32, {7, 9}, std::vector{DExpr::Var(34), DExpr::Var(35)}); + const Shape padding_value_shape = ShapeUtil::MakeShape(F32, {}); + PaddingConfig padding_config; + auto* dimension0 = padding_config.add_dimensions(); + dimension0->set_edge_padding_low(2); + dimension0->set_edge_padding_high(1); + dimension0->set_interior_padding(2); + auto* dimension1 = padding_config.add_dimensions(); + dimension1->set_edge_padding_low(3); + dimension1->set_edge_padding_high(4); + dimension1->set_interior_padding(1); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferPadShape(input_shape, padding_value_shape, + padding_config)); + + EXPECT_TRUE(ShapeUtil::Equal( + inferred_shape, ShapeUtil::MakeShape( + F32, {22, 24}, + std::vector{DExpr::Var(34) + 15, + DExpr::Var(35) + 15}))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), + DExpr::Var(34) + 15)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), + DExpr::Var(35) + 15)); +} + TEST_F(ShapeInferenceTest, Reverse) { const Shape input_shape = ShapeUtil::MakeShape(F32, {10, 25}); @@ -3020,6 +3542,20 @@ TEST_F(ShapeInferenceTest, Reverse) { ASSERT_TRUE(ShapeUtil::Equal(input_shape, *inferred_shape)); } +TEST_F(ShapeInferenceTest, ReversePreservesExpressions) { + const Shape input_shape = ShapeUtil::MakeShape( + F32, {10, 25, 7}, + std::vector{DExpr::Var(26), DExpr::Const(25), DExpr::Var(27)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReverseShape(input_shape, {0, 2})); + + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(26))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(25))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(2), DExpr::Var(27))); +} + TEST_F(ShapeInferenceTest, ReverseInvalidDimension) { const Shape input_shape = ShapeUtil::MakeShape(F32, {10, 25}); @@ -3094,6 +3630,21 @@ TEST_F(ShapeInferenceTest, Transpose) { *inferred_shape_and_status)); } +TEST_F(ShapeInferenceTest, TransposePermutesExpressions) { + const Shape a_shape = ShapeUtil::MakeShape( + F32, {2, 3, 4, 5}, + std::vector{DExpr::Var(28), DExpr::Const(3), DExpr::Var(29), + DExpr::Const(5)}); + TF_ASSERT_OK_AND_ASSIGN(const Shape inferred_shape, + ShapeInference::InferTransposeShape( + a_shape, {1, 2, 3, 0})); + + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Const(3))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Var(29))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(2), DExpr::Const(5))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(3), DExpr::Var(28))); +} + TEST_F(ShapeInferenceTest, Rank1Transpose) { const Shape a_shape = ShapeUtil::MakeShape(F32, {5}); const absl::StatusOr inferred_shape_and_status = @@ -3365,6 +3916,25 @@ TEST_F(ShapeInferenceTest, GoodTopK) { ShapeUtil::MakeShape(S32, {3, 4, 2})}))); } +TEST_F(ShapeInferenceTest, TopKPreservesLeadingExpressions) { + const Shape input = ShapeUtil::MakeShape( + F32, {3, 4, 5}, + std::vector{DExpr::Var(7), DExpr::Const(4), DExpr::Var(8)}); + TF_ASSERT_OK_AND_ASSIGN(const Shape inferred_shape, + ShapeInference::InferTopKShape(input, /*k=*/2)); + + ASSERT_TRUE(inferred_shape.IsTuple()); + ASSERT_EQ(inferred_shape.tuple_shapes_size(), 2); + const Shape& values = inferred_shape.tuple_shapes(0); + const Shape& indices = inferred_shape.tuple_shapes(1); + EXPECT_TRUE(DynExpr::equal(values.expressions(0), DExpr::Var(7))); + EXPECT_TRUE(DynExpr::equal(values.expressions(1), DExpr::Const(4))); + EXPECT_TRUE(DynExpr::equal(values.expressions(2), DExpr::Const(2))); + EXPECT_TRUE(DynExpr::equal(indices.expressions(0), DExpr::Var(7))); + EXPECT_TRUE(DynExpr::equal(indices.expressions(1), DExpr::Const(4))); + EXPECT_TRUE(DynExpr::equal(indices.expressions(2), DExpr::Const(2))); +} + TEST_F(ShapeInferenceTest, FailTopKLargeK) { const Shape input = ShapeUtil::MakeShape(F32, {3, 4, 5}); const absl::StatusOr statusor = @@ -3566,6 +4136,30 @@ TEST_F(GatherShapeInferenceTest, DynamicIndices) { << ShapeUtil::HumanString(gather_shape); } +TEST_F(GatherShapeInferenceTest, GatherPreservesIndexAndSliceExpressions) { + const Shape input = ShapeUtil::MakeShape( + F32, {3, 2, 2}, + std::vector{DExpr::Const(3), DExpr::Var(22), DExpr::Const(2)}); + const Shape indices = ShapeUtil::MakeShape( + S64, {3, 4, 2}, std::vector{false, true, false}, + std::vector{DExpr::Const(3), DExpr::Var(23), DExpr::Const(2)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape gather_shape, + ShapeInference::InferGatherShape( + input, indices, + HloGatherInstruction::MakeGatherDimNumbers( + /*offset_dims=*/{2, 3}, + /*collapsed_slice_dims=*/{0}, + /*start_index_map=*/{0, 1}, + /*index_vector_dim=*/2), + /*slice_sizes=*/{1, 2, 2})); + EXPECT_TRUE(DynExpr::equal(gather_shape.expressions(0), DExpr::Const(3))); + EXPECT_TRUE(DynExpr::equal(gather_shape.expressions(1), DExpr::Var(23))); + EXPECT_TRUE(DynExpr::equal(gather_shape.expressions(2), DExpr::Var(22))); + EXPECT_TRUE(DynExpr::equal(gather_shape.expressions(3), DExpr::Const(2))); +} + TEST_F(GatherShapeInferenceTest, NonDefaultGatherIndicesLeafDim_A) { TF_ASSERT_OK_AND_ASSIGN( const Shape gather_shape, @@ -4734,6 +5328,20 @@ TEST_F(ShapeInferenceTest, UnboundedAllToAll) { << " expected: " << ShapeUtil::HumanString(expected); } +TEST_F(ShapeInferenceTest, AllToAllPreservesExpressions) { + const Shape operand = ShapeUtil::MakeShape( + F32, {12, 5}, + std::vector{DExpr::Var(14), DExpr::Const(5)}); + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferAllToAllShape(/*shape=*/operand, + /*split_dimension=*/0, + /*concat_dimension=*/0, + /*split_count=*/3)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(14))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(5))); +} + TEST_F(ShapeInferenceTest, UnboundedAllToAllTupleUnsupported) { TF_ASSERT_OK_AND_ASSIGN(const Shape operand, ParseShape("f32[?, 10]")); TF_ASSERT_OK_AND_ASSIGN(const Shape expected, @@ -4813,6 +5421,30 @@ TEST_F(ShapeInferenceTest, UnboundedBatchNormGrad) { << " expected: " << ShapeUtil::HumanString(expected_tuple_shape); } +TEST_F(ShapeInferenceTest, BatchNormGradPreservesFeatureExpression) { + const Shape operand = ShapeUtil::MakeShape( + F32, {5, 7, 11}, std::vector{false, true, false}, + std::vector{DExpr::Const(5), DExpr::Var(12), DExpr::Const(11)}); + const Shape scale = + ShapeUtil::MakeShape(F32, {7}, std::vector{DExpr::Var(12)}); + const Shape mean = + ShapeUtil::MakeShape(F32, {7}, std::vector{DExpr::Var(12)}); + const Shape variance = + ShapeUtil::MakeShape(F32, {7}, std::vector{DExpr::Var(12)}); + const Shape output_grad = operand; + + TF_ASSERT_OK_AND_ASSIGN(const Shape inferred_shape, + ShapeInference::InferBatchNormGradShape( + operand, scale, mean, variance, output_grad, 1)); + ASSERT_TRUE(inferred_shape.IsTuple()); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(0).expressions(1), DExpr::Var(12))); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(1).expressions(0), DExpr::Var(12))); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(2).expressions(0), DExpr::Var(12))); +} + TEST_F(ShapeInferenceTest, UnboundedBatchNormInference) { TF_ASSERT_OK_AND_ASSIGN(const Shape operand, ParseShape("f32[?, ?, 7]")); TF_ASSERT_OK_AND_ASSIGN(const Shape scale, ParseShape("f32[5]")); @@ -4845,6 +5477,27 @@ TEST_F(ShapeInferenceTest, UnboundedBatchNormTraining) { << " expected: " << ShapeUtil::HumanString(expected_tuple_shape); } +TEST_F(ShapeInferenceTest, BatchNormTrainingPreservesFeatureExpression) { + const Shape operand = ShapeUtil::MakeShape( + F32, {5, 7, 11}, std::vector{false, true, false}, + std::vector{DExpr::Const(5), DExpr::Var(13), DExpr::Const(11)}); + const Shape scale = + ShapeUtil::MakeShape(F32, {7}, std::vector{DExpr::Var(13)}); + const Shape offset = + ShapeUtil::MakeShape(F32, {7}, std::vector{DExpr::Var(13)}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferBatchNormTrainingShape(operand, scale, offset, 1)); + ASSERT_TRUE(inferred_shape.IsTuple()); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(0).expressions(1), DExpr::Var(13))); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(1).expressions(0), DExpr::Var(13))); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(2).expressions(0), DExpr::Var(13))); +} + TEST_F(ShapeInferenceTest, UnboundedBroadcastUnsupportedOperand) { TF_ASSERT_OK_AND_ASSIGN(const Shape operand, ParseShape("f32[<=2, ?]")); TF_ASSERT_OK_AND_ASSIGN(const Shape expected, ParseShape("f32[1, <=2, ?]")); @@ -5234,6 +5887,29 @@ TEST_F(ShapeInferenceTest, UnboundedDotGeneral) { << " expected: " << ShapeUtil::HumanString(expected); } +TEST_F(ShapeInferenceTest, DotGeneralPreservesBatchExpression) { + const Shape lhs = ShapeUtil::MakeShape( + F32, {2, 3, 5}, std::vector{true, false, false}, + {DExpr::Var(1), DExpr::Const(3), DExpr::Const(5)}); + const Shape rhs = ShapeUtil::MakeShape( + F32, {2, 5, 7}, std::vector{true, false, false}, + {DExpr::Var(1), DExpr::Const(5), DExpr::Const(7)}); + + DotDimensionNumbers dnums; + dnums.add_lhs_batch_dimensions(0); + dnums.add_rhs_batch_dimensions(0); + dnums.add_lhs_contracting_dimensions(2); + dnums.add_rhs_contracting_dimensions(1); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferDotOpShape(lhs, rhs, dnums, + /*preferred_element_type=*/std::nullopt)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(1))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(3))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(2), DExpr::Const(7))); +} + TEST_F(ShapeInferenceTest, UnboundedDynamicSlice) { TF_ASSERT_OK_AND_ASSIGN(const Shape operand, ParseShape("f32[?, 10]")); TF_ASSERT_OK_AND_ASSIGN(const Shape start_index, ParseShape("s32[]")); @@ -5249,6 +5925,23 @@ TEST_F(ShapeInferenceTest, UnboundedDynamicSlice) { << " expected: " << ShapeUtil::HumanString(expected); } +TEST_F(ShapeInferenceTest, DynamicSliceUsesProvidedExpressions) { + const Shape operand = ShapeUtil::MakeShape( + F32, {9, 10}, + std::vector{DExpr::Var(15), DExpr::Const(10)}); + const Shape start_index = ShapeUtil::MakeShape(S32, {}); + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferDynamicSliceShape( + operand, /*start_index_shapes=*/{start_index, start_index}, + /*slice_sizes=*/{4, 10}, + /*slice_exprs=*/{DExpr::Var(15) / 2, DExpr::Const(10)}, + /*allow_scalar_indices=*/true)); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(15) / 2)); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(10))); +} + TEST_F(ShapeInferenceTest, UnboundedDynamicUpdateSlice) { TF_ASSERT_OK_AND_ASSIGN(const Shape operand, ParseShape("f32[?, 10]")); TF_ASSERT_OK_AND_ASSIGN(const Shape update, ParseShape("f32[?, 5]")); @@ -5264,6 +5957,42 @@ TEST_F(ShapeInferenceTest, UnboundedDynamicUpdateSlice) { << " expected: " << ShapeUtil::HumanString(expected); } +TEST_F(ShapeInferenceTest, DynamicUpdateSlicePreservesOperandExpressions) { + const Shape operand = ShapeUtil::MakeShape( + F32, {12, 10}, + std::vector{DExpr::Var(30), DExpr::Const(10)}); + const Shape update = ShapeUtil::MakeShape( + F32, {4, 10}, + std::vector{DExpr::Const(4), DExpr::Const(10)}); + const Shape start_index = ShapeUtil::MakeShape(S32, {}); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferDynamicUpdateSliceShape( + operand, update, /*start_index_shapes=*/{start_index, start_index}, + /*allow_scalar_indices=*/true)); + + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(30))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Const(10))); +} + +TEST_F(ShapeInferenceTest, DynamicReshapePreservesExpressions) { + const Shape operand = ShapeUtil::MakeShape( + F32, {5, 4, 8}, + std::vector{false, false, false}, + std::vector{DExpr::Var(21), DExpr::Const(4), DExpr::Const(8)}); + const Shape dim_size = ShapeUtil::MakeShape(S32, {}); + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferDynamicReshapeShape( + operand, /*dim_size_shapes=*/{&dim_size}, + /*new_size_bounds=*/{160}, + /*dims_are_dynamic=*/{false}, + /*expressions=*/{DExpr::Var(21) * DExpr::Const(32)})); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), + DExpr::Var(21) * DExpr::Const(32))); +} + TEST_F(ShapeInferenceTest, UnboundedFftWithFFT) { TF_ASSERT_OK_AND_ASSIGN(const Shape operand, ParseShape("c64[2, <=5, ?]")); const std::vector fft_length = {5, 10}; @@ -5499,6 +6228,56 @@ TEST_F(ShapeInferenceTest, UnboundedReduce) { << " expected: " << ShapeUtil::HumanString(expected); } +TEST_F(ShapeInferenceTest, ReducePreservesRemainingExpressions) { + const Shape input = ShapeUtil::MakeShape( + F32, {5, 7, 11}, + std::vector{DExpr::Var(9), DExpr::Const(7), DExpr::Var(10)}); + ProgramShape to_apply = + ShapeUtil::MakeProgramShape({f32_, f32_}, f32_); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReduceShape({&input, &f32_}, {1}, to_apply)); + + EXPECT_TRUE(ShapeUtil::Equal( + inferred_shape, ShapeUtil::MakeShape( + F32, {5, 11}, + std::vector{DExpr::Var(9), DExpr::Var(10)}))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(0), DExpr::Var(9))); + EXPECT_TRUE(DynExpr::equal(inferred_shape.expressions(1), DExpr::Var(10))); +} + +TEST_F(ShapeInferenceTest, ReduceTupleOutputsPreserveRemainingExpressions) { + const Shape input0 = ShapeUtil::MakeShape( + F32, {5, 7, 11}, + std::vector{DExpr::Var(31), DExpr::Const(7), DExpr::Var(32)}); + const Shape input1 = ShapeUtil::MakeShape( + S32, {5, 7, 11}, + std::vector{DExpr::Var(31), DExpr::Const(7), DExpr::Var(32)}); + ProgramShape to_apply = ShapeUtil::MakeProgramShape( + {f32_, s32_, f32_, s32_}, ShapeUtil::MakeTupleShape({f32_, s32_})); + + TF_ASSERT_OK_AND_ASSIGN( + const Shape inferred_shape, + ShapeInference::InferReduceShape( + {&input0, &input1, &f32_, &s32_}, {1}, to_apply)); + + ASSERT_TRUE(inferred_shape.IsTuple()); + ASSERT_EQ(inferred_shape.tuple_shapes_size(), 2); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(0).expressions(0), + DExpr::Var(31))); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(0).expressions(1), + DExpr::Var(32))); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(1).expressions(0), + DExpr::Var(31))); + EXPECT_TRUE( + DynExpr::equal(inferred_shape.tuple_shapes(1).expressions(1), + DExpr::Var(32))); +} + TEST_F(ShapeInferenceTest, UnboundedReduceInvalidReduceDimension) { TF_ASSERT_OK_AND_ASSIGN(const Shape input0, ParseShape("f32[7, 5]")); TF_ASSERT_OK_AND_ASSIGN(const Shape input1, ParseShape("f32[?, 5]")); @@ -5723,6 +6502,47 @@ TEST_F(ShapeInferenceTest, UnboundedSelectAndScatter) { << " expected: " << ShapeUtil::HumanString(expected); } +TEST_F(ShapeInferenceTest, SelectAndScatterPreservesOperandExpressions) { + const Shape operand = ShapeUtil::MakeShape( + F32, {11, 10}, std::vector{DExpr::Var(36), DExpr::Const(10)}); + const Shape source = ShapeUtil::MakeShape( + F32, {5, 10}, std::vector{(((DExpr::Var(36) - 1) / 2) + 1).simplify(), + DExpr::Const(10)}); + const Shape init_value = ShapeUtil::MakeShape(F32, {}); + + Window window; + WindowDimension dim0; + dim0.set_base_dilation(1); + dim0.set_size(3); + dim0.set_stride(2); + dim0.set_padding_low(0); + dim0.set_padding_high(1); + dim0.set_window_dilation(1); + + WindowDimension dim1; + dim1.set_base_dilation(1); + dim1.set_size(1); + dim1.set_stride(1); + dim1.set_padding_low(0); + dim1.set_padding_high(0); + dim1.set_window_dilation(1); + + *window.add_dimensions() = dim0; + *window.add_dimensions() = dim1; + + TF_ASSERT_OK_AND_ASSIGN( + const Shape result, + ShapeInference::InferSelectAndScatterShape( + operand, + /*select_shape=*/ShapeUtil::MakeProgramShape({f32_, f32_}, pred_), + window, source, init_value, + /*scatter_shape=*/ + ShapeUtil::MakeProgramShape({f32_, f32_}, f32_))); + + EXPECT_TRUE(DynExpr::equal(result.expressions(0), DExpr::Var(36))); + EXPECT_TRUE(DynExpr::equal(result.expressions(1), DExpr::Const(10))); +} + TEST_P(UnboundedBinaryOpShapeInferenceTest, UnboundedShiftLeft) { TF_ASSERT_OK_AND_ASSIGN(const Shape lhs, ParseShape(GetParam().lhs)); TF_ASSERT_OK_AND_ASSIGN(const Shape rhs, ParseShape(GetParam().rhs)); diff --git a/third_party/xla/xla/shape_expr.cc b/third_party/xla/xla/shape_expr.cc index 21dad324e8d44f..5beab5cedbca75 100644 --- a/third_party/xla/xla/shape_expr.cc +++ b/third_party/xla/xla/shape_expr.cc @@ -317,10 +317,24 @@ std::unique_ptr SimplifyFallback(const DynExpr* expr) { } Constant* l = AsConstant(lhs.get()); Constant* r = AsConstant(rhs.get()); + if (*lhs == *rhs) { + return std::make_unique(1); + } if (l && l->get_val() == 0 && r && r->get_val() != 0) { return std::make_unique(0); } if (r && r->get_val() == 1) return lhs; + if (lhs->kind() == DExpr::Kind::kMul) { + auto* mul = static_cast(lhs.get()); + auto lhs_l = std::unique_ptr(mul->get_lhs()->s()); + auto lhs_r = std::unique_ptr(mul->get_rhs()->s()); + if (*lhs_l == *rhs) { + return lhs_r; + } + if (*lhs_r == *rhs) { + return lhs_l; + } + } if (l && r && r->get_val() != 0) { int64_t numerator = l->get_val(); int64_t denominator = r->get_val(); @@ -391,6 +405,13 @@ bool DynExpr::equal(DynExpr* expr1, DynExpr* expr2) { auto e1 = std::unique_ptr(expr1->s()); auto e2 = std::unique_ptr(expr2->s()); if (e1 == nullptr || e2 == nullptr) return false; + auto a1 = ToCanonicalAffine(e1.get()); + auto a2 = ToCanonicalAffine(e2.get()); + if (a1.has_value() && a2.has_value()) { + return a1->denominator == a2->denominator && + a1->constant == a2->constant && + a1->coefficients == a2->coefficients; + } if (e1->kind() == DExpr::Kind::kConstant && e2->kind() == DExpr::Kind::kConstant) { return static_cast(e1.get())->get_val() == From 5e677f77c9e80da26568d8432c92898f59278e3f Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Wed, 29 Jul 2026 12:40:59 +0100 Subject: [PATCH 02/15] Support max and conditional dynamic shape expressions (#4) --- .../jit/encapsulate_subgraphs_pass.cc | 11 ++ tensorflow/compiler/jit/kernels/xla_ops.cc | 33 ++++ .../compiler/jit/mark_for_compilation_pass.cc | 44 +++++ .../compiler/tf2xla/kernels/const_op.cc | 26 +++ .../compiler/tf2xla/kernels/reshape_op.cc | 8 +- .../compiler/tf2xla/kernels/sequence_ops.cc | 22 +-- .../tf2xla/kernels/strided_slice_op.cc | 14 +- tensorflow/core/framework/tensor_shape.cc | 39 ++++ tensorflow/core/framework/tensor_shape.proto | 21 +- .../core/framework/tensor_shape_expr.cc | 88 +++++++++ tensorflow/core/framework/tensor_shape_expr.h | 86 ++++++++ .../core/grappler/costs/graph_properties.cc | 3 + third_party/xla/xla/hlo/ir/hlo_instruction.cc | 16 ++ .../xla/xla/service/llvm_ir/llvm_util.cc | 26 +++ .../xla/xla/service/shape_inference.cc | 23 ++- .../xla/xla/service/shape_inference_test.cc | 52 +++++ third_party/xla/xla/shape_expr.cc | 100 +++++++++- third_party/xla/xla/shape_expr.h | 185 +++++++++++++++++- third_party/xla/xla/shape_test.cc | 39 +++- third_party/xla/xla/xla_data.proto | 21 +- 20 files changed, 816 insertions(+), 41 deletions(-) diff --git a/tensorflow/compiler/jit/encapsulate_subgraphs_pass.cc b/tensorflow/compiler/jit/encapsulate_subgraphs_pass.cc index a3d7b0e71783d3..080bef205441f2 100644 --- a/tensorflow/compiler/jit/encapsulate_subgraphs_pass.cc +++ b/tensorflow/compiler/jit/encapsulate_subgraphs_pass.cc @@ -139,6 +139,17 @@ std::string ExprProtoToString(const ExpressionProto& e) { case ExpressionProto::kDivNode: return absl::StrCat("(", ExprProtoToString(e.div_node().lhs()), " / ", ExprProtoToString(e.div_node().rhs()), ")"); + case ExpressionProto::kMaxNode: + return absl::StrCat("max(", ExprProtoToString(e.max_node().lhs()), ", ", + ExprProtoToString(e.max_node().rhs()), ")"); + case ExpressionProto::kGtNode: + return absl::StrCat("(", ExprProtoToString(e.gt_node().lhs()), " > ", + ExprProtoToString(e.gt_node().rhs()), ")"); + case ExpressionProto::kSelectNode: + return absl::StrCat("select(", ExprProtoToString(e.select_node().pred()), + ", ", ExprProtoToString(e.select_node().on_true()), + ", ", ExprProtoToString(e.select_node().on_false()), + ")"); default: return ""; } diff --git a/tensorflow/compiler/jit/kernels/xla_ops.cc b/tensorflow/compiler/jit/kernels/xla_ops.cc index 2fb97d66b3ac63..1a84226b2133ab 100644 --- a/tensorflow/compiler/jit/kernels/xla_ops.cc +++ b/tensorflow/compiler/jit/kernels/xla_ops.cc @@ -414,6 +414,23 @@ std::unique_ptr ExprFromProto(const ExpressionProto& proto) { auto rhs = ExprFromProto(proto.div_node().rhs()); return std::make_unique(lhs.release(), rhs.release()); } + case ExpressionProto::kMaxNode: { + auto lhs = ExprFromProto(proto.max_node().lhs()); + auto rhs = ExprFromProto(proto.max_node().rhs()); + return std::make_unique(lhs.release(), rhs.release()); + } + case ExpressionProto::kGtNode: { + auto lhs = ExprFromProto(proto.gt_node().lhs()); + auto rhs = ExprFromProto(proto.gt_node().rhs()); + return std::make_unique(lhs.release(), rhs.release()); + } + case ExpressionProto::kSelectNode: { + auto pred = ExprFromProto(proto.select_node().pred()); + auto on_true = ExprFromProto(proto.select_node().on_true()); + auto on_false = ExprFromProto(proto.select_node().on_false()); + return std::make_unique(pred.release(), on_true.release(), + on_false.release()); + } case ExpressionProto::NODE_TYPE_NOT_SET: default: return nullptr; @@ -445,6 +462,22 @@ static xla::DExpr DimExprToDExpr(const DimExpr* e) { auto* ee = static_cast(e); return DimExprToDExpr(ee->lhs()) / DimExprToDExpr(ee->rhs()); } + case DimExpr::Kind::kMax: { + auto* ee = static_cast(e); + return xla::DExpr::Max(DimExprToDExpr(ee->lhs()), + DimExprToDExpr(ee->rhs())); + } + case DimExpr::Kind::kGt: { + auto* ee = static_cast(e); + return xla::DExpr::Gt(DimExprToDExpr(ee->lhs()), + DimExprToDExpr(ee->rhs())); + } + case DimExpr::Kind::kSelect: { + auto* ee = static_cast(e); + return xla::DExpr::Select(DimExprToDExpr(ee->pred()), + DimExprToDExpr(ee->on_true()), + DimExprToDExpr(ee->on_false())); + } } return xla::DExpr::Unknown(); } diff --git a/tensorflow/compiler/jit/mark_for_compilation_pass.cc b/tensorflow/compiler/jit/mark_for_compilation_pass.cc index 566bab23a11867..1b320d1b332d90 100644 --- a/tensorflow/compiler/jit/mark_for_compilation_pass.cc +++ b/tensorflow/compiler/jit/mark_for_compilation_pass.cc @@ -714,6 +714,17 @@ std::string ExprProtoToString(const ExpressionProto& e) { case ExpressionProto::kDivNode: return absl::StrCat("(", ExprProtoToString(e.div_node().lhs()), " / ", ExprProtoToString(e.div_node().rhs()), ")"); + case ExpressionProto::kMaxNode: + return absl::StrCat("max(", ExprProtoToString(e.max_node().lhs()), ", ", + ExprProtoToString(e.max_node().rhs()), ")"); + case ExpressionProto::kGtNode: + return absl::StrCat("(", ExprProtoToString(e.gt_node().lhs()), " > ", + ExprProtoToString(e.gt_node().rhs()), ")"); + case ExpressionProto::kSelectNode: + return absl::StrCat("select(", ExprProtoToString(e.select_node().pred()), + ", ", ExprProtoToString(e.select_node().on_true()), + ", ", ExprProtoToString(e.select_node().on_false()), + ")"); default: return ""; } @@ -747,6 +758,23 @@ std::unique_ptr ExprFromProto(const ExpressionProto& proto) { auto rhs = ExprFromProto(proto.div_node().rhs()); return std::make_unique(lhs.release(), rhs.release()); } + case ExpressionProto::kMaxNode: { + auto lhs = ExprFromProto(proto.max_node().lhs()); + auto rhs = ExprFromProto(proto.max_node().rhs()); + return std::make_unique(lhs.release(), rhs.release()); + } + case ExpressionProto::kGtNode: { + auto lhs = ExprFromProto(proto.gt_node().lhs()); + auto rhs = ExprFromProto(proto.gt_node().rhs()); + return std::make_unique(lhs.release(), rhs.release()); + } + case ExpressionProto::kSelectNode: { + auto pred = ExprFromProto(proto.select_node().pred()); + auto on_true = ExprFromProto(proto.select_node().on_true()); + auto on_false = ExprFromProto(proto.select_node().on_false()); + return std::make_unique(pred.release(), on_true.release(), + on_false.release()); + } case ExpressionProto::NODE_TYPE_NOT_SET: default: return nullptr; @@ -779,6 +807,22 @@ static xla::DExpr DimExprToDExpr(const DimExpr* e) { auto* ee = static_cast(e); return DimExprToDExpr(ee->lhs()) / DimExprToDExpr(ee->rhs()); } + case DimExpr::Kind::kMax: { + auto* ee = static_cast(e); + return xla::DExpr::Max(DimExprToDExpr(ee->lhs()), + DimExprToDExpr(ee->rhs())); + } + case DimExpr::Kind::kGt: { + auto* ee = static_cast(e); + return xla::DExpr::Gt(DimExprToDExpr(ee->lhs()), + DimExprToDExpr(ee->rhs())); + } + case DimExpr::Kind::kSelect: { + auto* ee = static_cast(e); + return xla::DExpr::Select(DimExprToDExpr(ee->pred()), + DimExprToDExpr(ee->on_true()), + DimExprToDExpr(ee->on_false())); + } } return xla::DExpr(); } diff --git a/tensorflow/compiler/tf2xla/kernels/const_op.cc b/tensorflow/compiler/tf2xla/kernels/const_op.cc index a911f28246da77..a67c25be8f14f2 100644 --- a/tensorflow/compiler/tf2xla/kernels/const_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/const_op.cc @@ -123,6 +123,16 @@ bool IsDynamicExpressionProto(const ExpressionProto& proto) { case ExpressionProto::kDivNode: return IsDynamicExpressionProto(proto.div_node().lhs()) || IsDynamicExpressionProto(proto.div_node().rhs()); + case ExpressionProto::kMaxNode: + return IsDynamicExpressionProto(proto.max_node().lhs()) || + IsDynamicExpressionProto(proto.max_node().rhs()); + case ExpressionProto::kGtNode: + return IsDynamicExpressionProto(proto.gt_node().lhs()) || + IsDynamicExpressionProto(proto.gt_node().rhs()); + case ExpressionProto::kSelectNode: + return IsDynamicExpressionProto(proto.select_node().pred()) || + IsDynamicExpressionProto(proto.select_node().on_true()) || + IsDynamicExpressionProto(proto.select_node().on_false()); case ExpressionProto::kConstantValue: case ExpressionProto::NODE_TYPE_NOT_SET: return false; @@ -158,6 +168,22 @@ static xla::DExpr DimExprToDExpr(const DimExpr* e) { const auto* ee = static_cast(e); return DimExprToDExpr(ee->lhs()) / DimExprToDExpr(ee->rhs()); } + case DimExpr::Kind::kMax: { + const auto* ee = static_cast(e); + return xla::DExpr::Max(DimExprToDExpr(ee->lhs()), + DimExprToDExpr(ee->rhs())); + } + case DimExpr::Kind::kGt: { + auto* ee = static_cast(e); + return xla::DExpr::Gt(DimExprToDExpr(ee->lhs()), + DimExprToDExpr(ee->rhs())); + } + case DimExpr::Kind::kSelect: { + auto* ee = static_cast(e); + return xla::DExpr::Select(DimExprToDExpr(ee->pred()), + DimExprToDExpr(ee->on_true()), + DimExprToDExpr(ee->on_false())); + } } return xla::DExpr(); } diff --git a/tensorflow/compiler/tf2xla/kernels/reshape_op.cc b/tensorflow/compiler/tf2xla/kernels/reshape_op.cc index c019f42927bb9b..eb861e895ac03b 100644 --- a/tensorflow/compiler/tf2xla/kernels/reshape_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/reshape_op.cc @@ -164,11 +164,11 @@ class ReshapeOp : public XlaOpKernel { input, xla::Zero(ctx->builder(), input_xla_shape->element_type()), 0, 0, padded_input_num - input_num_elements); input_shape.set_dim(0, padded_input_num); - // This expression only approximates the padded size: the true value - // uses ceil(input_num_elements / product) * product, which we do not - // model symbolically here. + missing_expr = + (input_num_elements_expr + (product - 1)) / product; + missing_expr = missing_expr.simplify(); xla::DExpr padded_input_num_expr = - ((input_num_elements_expr / product_expr) * product_expr) + (missing_expr * xla::DExpr::Const(product)) .simplify(); input_shape.set_expression(0, padded_input_num_expr); } diff --git a/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc b/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc index 1ff48ac5d64a66..1e87b61d7fc19e 100644 --- a/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc +++ b/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc @@ -92,17 +92,17 @@ xla::DExpr BuildRangeSizeExpr(const XlaExpression& start_expr, xla::DExpr limit_symbol = GetScalarExpr(limit_expr, limit); xla::DExpr delta_symbol = GetScalarExpr(delta_expr, delta); - if (delta.Get({}) > 0) { - xla::DExpr diff = (limit_symbol - start_symbol).simplify(); - xla::DExpr adjusted = (diff - 1).simplify(); - xla::DExpr quotient = (adjusted / delta_symbol).simplify(); - return (quotient + 1).simplify(); - } - xla::DExpr step_symbol = (xla::DExpr::Const(0) - delta_symbol).simplify(); - xla::DExpr diff = (start_symbol - limit_symbol).simplify(); - xla::DExpr adjusted = (diff - 1).simplify(); - xla::DExpr quotient = (adjusted / step_symbol).simplify(); - return (quotient + 1).simplify(); + xla::DExpr positive_diff = (limit_symbol - start_symbol).simplify(); + xla::DExpr positive_size = + (((positive_diff - 1) / delta_symbol) + 1).simplify(); + xla::DExpr negative_step = + (xla::DExpr::Const(0) - delta_symbol).simplify(); + xla::DExpr negative_diff = (start_symbol - limit_symbol).simplify(); + xla::DExpr negative_size = + (((negative_diff - 1) / negative_step) + 1).simplify(); + return xla::DExpr::Select(xla::DExpr::Gt(delta_symbol, xla::DExpr::Const(0)), + positive_size, negative_size) + .simplify(); } // The type-specific part of the implementation of Range. diff --git a/tensorflow/compiler/tf2xla/kernels/strided_slice_op.cc b/tensorflow/compiler/tf2xla/kernels/strided_slice_op.cc index 9189a83d3f7f69..d776bbf68e8525 100644 --- a/tensorflow/compiler/tf2xla/kernels/strided_slice_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/strided_slice_op.cc @@ -348,8 +348,8 @@ class StridedSliceOp : public XlaOpKernel { slice_begin.push_back(begin[i]); slice_begin_expr.push_back(begin_expr[i]); slice_end.push_back(std::max(end[i], begin[i])); - slice_end_expr.push_back((end[i] > begin[i]) ? end_expr[i] - : begin_expr[i]); + slice_end_expr.push_back( + xla::DExpr::Max(end_expr[i], begin_expr[i]).simplify()); slice_strides.push_back(strides[i]); } else { // Negative stride: swap begin and end, add 1 because the interval @@ -361,11 +361,11 @@ class StridedSliceOp : public XlaOpKernel { slice_end.push_back(std::max(input_shape.dim_size(i) - end[i] - 1, input_shape.dim_size(i) - begin[i] - 1)); slice_end_expr.push_back( - (end[i] < begin[i]) - ? (input_expr - end_expr[i] - xla::DExpr::Const(1)) - .simplify() - : (input_expr - begin_expr[i] - xla::DExpr::Const(1)) - .simplify()); + xla::DExpr::Max( + (input_expr - end_expr[i] - xla::DExpr::Const(1)).simplify(), + (input_expr - begin_expr[i] - xla::DExpr::Const(1)) + .simplify()) + .simplify()); slice_strides.push_back(-strides[i]); dimensions_to_reverse.push_back(i); } diff --git a/tensorflow/core/framework/tensor_shape.cc b/tensorflow/core/framework/tensor_shape.cc index 0ceeb8c9b169f6..f42571ea10cca2 100644 --- a/tensorflow/core/framework/tensor_shape.cc +++ b/tensorflow/core/framework/tensor_shape.cc @@ -52,6 +52,21 @@ xla::DExpr DExprFromProto(const ExpressionProto& proto) { const auto& div = proto.div_node(); return DExprFromProto(div.lhs()) / DExprFromProto(div.rhs()); } + case ExpressionProto::kMaxNode: { + const auto& max = proto.max_node(); + return xla::DExpr::Max(DExprFromProto(max.lhs()), + DExprFromProto(max.rhs())); + } + case ExpressionProto::kGtNode: { + const auto& gt = proto.gt_node(); + return xla::DExpr::Gt(DExprFromProto(gt.lhs()), DExprFromProto(gt.rhs())); + } + case ExpressionProto::kSelectNode: { + const auto& select = proto.select_node(); + return xla::DExpr::Select(DExprFromProto(select.pred()), + DExprFromProto(select.on_true()), + DExprFromProto(select.on_false())); + } case ExpressionProto::NODE_TYPE_NOT_SET: default: return xla::DExpr::Unknown(xla::kMissingExpressionSentinel); @@ -106,6 +121,30 @@ void ExprToProto(const xla::DExpr& expr, ExpressionProto* proto) { div->mutable_rhs()); return; } + case xla::DExpr::Kind::kMax: { + auto* max = proto->mutable_max_node(); + const auto& node = static_cast(*expr.get()); + ExprToProto(xla::DExpr(node.get_lhs()->clone()), max->mutable_lhs()); + ExprToProto(xla::DExpr(node.get_rhs()->clone()), max->mutable_rhs()); + return; + } + case xla::DExpr::Kind::kGt: { + auto* gt = proto->mutable_gt_node(); + const auto& node = static_cast(*expr.get()); + ExprToProto(xla::DExpr(node.get_lhs()->clone()), gt->mutable_lhs()); + ExprToProto(xla::DExpr(node.get_rhs()->clone()), gt->mutable_rhs()); + return; + } + case xla::DExpr::Kind::kSelect: { + auto* select = proto->mutable_select_node(); + const auto& node = static_cast(*expr.get()); + ExprToProto(xla::DExpr(node.get_pred()->clone()), select->mutable_pred()); + ExprToProto(xla::DExpr(node.get_on_true()->clone()), + select->mutable_on_true()); + ExprToProto(xla::DExpr(node.get_on_false()->clone()), + select->mutable_on_false()); + return; + } } } diff --git a/tensorflow/core/framework/tensor_shape.proto b/tensorflow/core/framework/tensor_shape.proto index f69b4228a7fb31..efddff41dc64ba 100644 --- a/tensorflow/core/framework/tensor_shape.proto +++ b/tensorflow/core/framework/tensor_shape.proto @@ -60,6 +60,9 @@ message ExpressionProto { SubNode sub_node = 4; // exp - exp MulNode mul_node = 5; // exp * exp DivNode div_node = 6; // exp / exp + MaxNode max_node = 7; // max(exp, exp) + GtNode gt_node = 8; // exp > exp + SelectNode select_node = 9; // select(pred, on_true, on_false) } } @@ -81,4 +84,20 @@ message MulNode { message DivNode { ExpressionProto lhs = 1; ExpressionProto rhs = 2; -} \ No newline at end of file +} + +message MaxNode { + ExpressionProto lhs = 1; + ExpressionProto rhs = 2; +} + +message GtNode { + ExpressionProto lhs = 1; + ExpressionProto rhs = 2; +} + +message SelectNode { + ExpressionProto pred = 1; + ExpressionProto on_true = 2; + ExpressionProto on_false = 3; +} diff --git a/tensorflow/core/framework/tensor_shape_expr.cc b/tensorflow/core/framework/tensor_shape_expr.cc index 686992d4c9881a..98607d439c3289 100644 --- a/tensorflow/core/framework/tensor_shape_expr.cc +++ b/tensorflow/core/framework/tensor_shape_expr.cc @@ -54,6 +54,16 @@ bool IsDynamicDimExpr(const ExpressionProto& proto) { case ExpressionProto::kDivNode: return IsDynamicDimExpr(proto.div_node().lhs()) || IsDynamicDimExpr(proto.div_node().rhs()); + case ExpressionProto::kMaxNode: + return IsDynamicDimExpr(proto.max_node().lhs()) || + IsDynamicDimExpr(proto.max_node().rhs()); + case ExpressionProto::kGtNode: + return IsDynamicDimExpr(proto.gt_node().lhs()) || + IsDynamicDimExpr(proto.gt_node().rhs()); + case ExpressionProto::kSelectNode: + return IsDynamicDimExpr(proto.select_node().pred()) || + IsDynamicDimExpr(proto.select_node().on_true()) || + IsDynamicDimExpr(proto.select_node().on_false()); case ExpressionProto::kConstantValue: case ExpressionProto::NODE_TYPE_NOT_SET: return false; @@ -123,6 +133,27 @@ static bool EqualsImpl(const DimExpr* a, const DimExpr* b) { return EqualsImpl(ad->lhs(), bd->lhs()) && EqualsImpl(ad->rhs(), bd->rhs()); } + case DimExpr::Kind::kMax: { + auto* am = static_cast(a); + auto* bm = static_cast(b); + return (EqualsImpl(am->lhs(), bm->lhs()) && + EqualsImpl(am->rhs(), bm->rhs())) || + (EqualsImpl(am->lhs(), bm->rhs()) && + EqualsImpl(am->rhs(), bm->lhs())); + } + case DimExpr::Kind::kGt: { + auto* ag = static_cast(a); + auto* bg = static_cast(b); + return EqualsImpl(ag->lhs(), bg->lhs()) && + EqualsImpl(ag->rhs(), bg->rhs()); + } + case DimExpr::Kind::kSelect: { + auto* as = static_cast(a); + auto* bs = static_cast(b); + return EqualsImpl(as->pred(), bs->pred()) && + EqualsImpl(as->on_true(), bs->on_true()) && + EqualsImpl(as->on_false(), bs->on_false()); + } } return false; @@ -160,6 +191,23 @@ std::unique_ptr DimExpr::FromProto(const ExpressionProto& proto) { auto rhs = FromProto(proto.div_node().rhs()); return std::make_unique(lhs.release(), rhs.release()); } + case ExpressionProto::kMaxNode: { + auto lhs = FromProto(proto.max_node().lhs()); + auto rhs = FromProto(proto.max_node().rhs()); + return std::make_unique(lhs.release(), rhs.release()); + } + case ExpressionProto::kGtNode: { + auto lhs = FromProto(proto.gt_node().lhs()); + auto rhs = FromProto(proto.gt_node().rhs()); + return std::make_unique(lhs.release(), rhs.release()); + } + case ExpressionProto::kSelectNode: { + auto pred = FromProto(proto.select_node().pred()); + auto on_true = FromProto(proto.select_node().on_true()); + auto on_false = FromProto(proto.select_node().on_false()); + return std::make_unique(pred.release(), on_true.release(), + on_false.release()); + } case ExpressionProto::NODE_TYPE_NOT_SET: default: return nullptr; @@ -255,6 +303,46 @@ DimExpr* SimplifyExpr(DimExpr* expr, return own(std::make_unique(lhs, rhs)); } + case DimExpr::Kind::kMax: { + auto* max = static_cast(expr); + DimExpr* lhs = SimplifyExpr(max->lhs(), arena); + DimExpr* rhs = SimplifyExpr(max->rhs(), arena); + if (lhs->IsConstant() && rhs->IsConstant()) { + return own(DimExpr::Cons( + std::max(lhs->ConstantValue(), rhs->ConstantValue()))); + } + if (lhs->IsConstant() && lhs->ConstantValue() == 0 && + rhs->kind() == DimExpr::Kind::kVariable) { + return rhs; + } + if (rhs->IsConstant() && rhs->ConstantValue() == 0 && + lhs->kind() == DimExpr::Kind::kVariable) { + return lhs; + } + if (DimExpr::Equals(lhs, rhs)) return lhs; + return own(std::make_unique(lhs, rhs)); + } + case DimExpr::Kind::kGt: { + auto* gt = static_cast(expr); + DimExpr* lhs = SimplifyExpr(gt->lhs(), arena); + DimExpr* rhs = SimplifyExpr(gt->rhs(), arena); + if (lhs->IsConstant() && rhs->IsConstant()) { + return own(DimExpr::Cons(lhs->ConstantValue() > rhs->ConstantValue())); + } + if (DimExpr::Equals(lhs, rhs)) return own(DimExpr::Cons(0)); + return own(std::make_unique(lhs, rhs)); + } + case DimExpr::Kind::kSelect: { + auto* select = static_cast(expr); + DimExpr* pred = SimplifyExpr(select->pred(), arena); + DimExpr* on_true = SimplifyExpr(select->on_true(), arena); + DimExpr* on_false = SimplifyExpr(select->on_false(), arena); + if (DimExpr::Equals(on_true, on_false)) return on_true; + if (pred->IsConstant()) { + return pred->ConstantValue() != 0 ? on_true : on_false; + } + return own(std::make_unique(pred, on_true, on_false)); + } } return expr; diff --git a/tensorflow/core/framework/tensor_shape_expr.h b/tensorflow/core/framework/tensor_shape_expr.h index 0a2497c03b808e..4ee7b79c7a74fa 100644 --- a/tensorflow/core/framework/tensor_shape_expr.h +++ b/tensorflow/core/framework/tensor_shape_expr.h @@ -1,6 +1,7 @@ #ifndef TENSORFLOW_CORE_FRAMEWORK_TENSOR_SHAPE_EXPR_H_ #define TENSORFLOW_CORE_FRAMEWORK_TENSOR_SHAPE_EXPR_H_ +#include #include #include #include @@ -18,6 +19,9 @@ class ExprAdd; class ExprSub; class ExprMul; class ExprDiv; +class ExprMax; +class ExprGt; +class ExprSelect; // DimExpr: Base class for symbolic expressions representing dynamic dimension // sizes. These expressions form a DAG that tracks how unknown dimensions relate @@ -27,6 +31,9 @@ class ExprDiv; // - Var(sym_id): A symbolic variable representing an unknown dimension // - Const(k): A known constant value // - Add/Sub/Mul/Div(lhs, rhs): Binary arithmetic operations +// - Max(lhs, rhs): Maximum of two expressions +// - Gt(lhs, rhs): Integer predicate (1 if lhs > rhs, otherwise 0) +// - Select(pred, on_true, on_false): Conditional expression // // INVARIANT: An unknown dimension is not just -1, it is -1 + Var(sym). class DimExpr { @@ -38,6 +45,9 @@ class DimExpr { kSub, kMul, kDiv, + kMax, + kGt, + kSelect, }; virtual ~DimExpr() = default; @@ -209,6 +219,82 @@ class ExprDiv final : public DimExpr { DimExpr* rhs_; }; +class ExprMax final : public DimExpr { + public: + ExprMax(DimExpr* lhs, DimExpr* rhs) : lhs_(lhs), rhs_(rhs) {} + + Kind kind() const override { return Kind::kMax; } + void ToProto(ExpressionProto* proto) const override { + auto* max_msg = proto->mutable_max_node(); + lhs_->ToProto(max_msg->mutable_lhs()); + rhs_->ToProto(max_msg->mutable_rhs()); + } + bool IsConstant() const override { + return lhs_->IsConstant() && rhs_->IsConstant(); + } + int64_t ConstantValue() const override { + return std::max(lhs_->ConstantValue(), rhs_->ConstantValue()); + } + DimExpr* lhs() const { return lhs_; } + DimExpr* rhs() const { return rhs_; } + + private: + DimExpr* lhs_; + DimExpr* rhs_; +}; + +class ExprGt final : public DimExpr { + public: + ExprGt(DimExpr* lhs, DimExpr* rhs) : lhs_(lhs), rhs_(rhs) {} + Kind kind() const override { return Kind::kGt; } + void ToProto(ExpressionProto* proto) const override { + auto* msg = proto->mutable_gt_node(); + lhs_->ToProto(msg->mutable_lhs()); + rhs_->ToProto(msg->mutable_rhs()); + } + bool IsConstant() const override { + return lhs_->IsConstant() && rhs_->IsConstant(); + } + int64_t ConstantValue() const override { + return lhs_->ConstantValue() > rhs_->ConstantValue(); + } + DimExpr* lhs() const { return lhs_; } + DimExpr* rhs() const { return rhs_; } + + private: + DimExpr* lhs_; + DimExpr* rhs_; +}; + +class ExprSelect final : public DimExpr { + public: + ExprSelect(DimExpr* pred, DimExpr* on_true, DimExpr* on_false) + : pred_(pred), on_true_(on_true), on_false_(on_false) {} + Kind kind() const override { return Kind::kSelect; } + void ToProto(ExpressionProto* proto) const override { + auto* msg = proto->mutable_select_node(); + pred_->ToProto(msg->mutable_pred()); + on_true_->ToProto(msg->mutable_on_true()); + on_false_->ToProto(msg->mutable_on_false()); + } + bool IsConstant() const override { + return pred_->IsConstant() && on_true_->IsConstant() && + on_false_->IsConstant(); + } + int64_t ConstantValue() const override { + return pred_->ConstantValue() != 0 ? on_true_->ConstantValue() + : on_false_->ConstantValue(); + } + DimExpr* pred() const { return pred_; } + DimExpr* on_true() const { return on_true_; } + DimExpr* on_false() const { return on_false_; } + + private: + DimExpr* pred_; + DimExpr* on_true_; + DimExpr* on_false_; +}; + // Simplify an expression tree: constant folding and algebraic identities. // Returns a NEW expression (does not mutate input). // The arena parameter is used to allocate nodes that will be owned externally. diff --git a/tensorflow/core/grappler/costs/graph_properties.cc b/tensorflow/core/grappler/costs/graph_properties.cc index 774c628ce07c34..e7d02aaa8adfbf 100644 --- a/tensorflow/core/grappler/costs/graph_properties.cc +++ b/tensorflow/core/grappler/costs/graph_properties.cc @@ -2488,6 +2488,9 @@ class SymbolicShapeManager { case DimExpr::Kind::kSub: case DimExpr::Kind::kMul: case DimExpr::Kind::kDiv: + case DimExpr::Kind::kMax: + case DimExpr::Kind::kGt: + case DimExpr::Kind::kSelect: return true; default: return false; diff --git a/third_party/xla/xla/hlo/ir/hlo_instruction.cc b/third_party/xla/xla/hlo/ir/hlo_instruction.cc index a7afb175ebc4ce..86b7bb0facc892 100644 --- a/third_party/xla/xla/hlo/ir/hlo_instruction.cc +++ b/third_party/xla/xla/hlo/ir/hlo_instruction.cc @@ -135,6 +135,22 @@ DynExpr* DynExprFromProtoForPrint(const ExpressionProto& proto) { return new Div(DynExprFromProtoForPrint(div.lhs()), DynExprFromProtoForPrint(div.rhs())); } + case ExpressionProto::kMaxNode: { + const auto& max = proto.max_node(); + return new MaxExpr(DynExprFromProtoForPrint(max.lhs()), + DynExprFromProtoForPrint(max.rhs())); + } + case ExpressionProto::kGtNode: { + const auto& gt = proto.gt_node(); + return new GtExpr(DynExprFromProtoForPrint(gt.lhs()), + DynExprFromProtoForPrint(gt.rhs())); + } + case ExpressionProto::kSelectNode: { + const auto& select = proto.select_node(); + return new SelectExpr(DynExprFromProtoForPrint(select.pred()), + DynExprFromProtoForPrint(select.on_true()), + DynExprFromProtoForPrint(select.on_false())); + } case ExpressionProto::NODE_TYPE_NOT_SET: default: return nullptr; diff --git a/third_party/xla/xla/service/llvm_ir/llvm_util.cc b/third_party/xla/xla/service/llvm_ir/llvm_util.cc index 8611b7f2fbd68d..6218ab26849cc7 100644 --- a/third_party/xla/xla/service/llvm_ir/llvm_util.cc +++ b/third_party/xla/xla/service/llvm_ir/llvm_util.cc @@ -919,6 +919,32 @@ static llvm::Value* EmitExpressionImpl(llvm::IRBuilderBase* b, llvm::Value* v_rhs = EmitExpressionImpl(b, *sub_node->get_rhs()); return b->CreateSub(v_lhs, v_rhs, "sub_dims"); } + if (expr.kind() == DExpr::Kind::kMax) { + auto* max_node = static_cast(&expr); + llvm::Value* v_lhs = EmitExpressionImpl(b, *max_node->get_lhs()); + llvm::Value* v_rhs = EmitExpressionImpl(b, *max_node->get_rhs()); + llvm::Value* lhs_is_greater = + b->CreateICmpSGT(v_lhs, v_rhs, "max_dims_pred"); + return b->CreateSelect(lhs_is_greater, v_lhs, v_rhs, "max_dims"); + } + if (expr.kind() == DExpr::Kind::kGt) { + auto* gt_node = static_cast(&expr); + llvm::Value* v_lhs = EmitExpressionImpl(b, *gt_node->get_lhs()); + llvm::Value* v_rhs = EmitExpressionImpl(b, *gt_node->get_rhs()); + llvm::Value* pred = b->CreateICmpSGT(v_lhs, v_rhs, "gt_dims_pred"); + return b->CreateZExt(pred, i64Type, "gt_dims"); + } + if (expr.kind() == DExpr::Kind::kSelect) { + auto* select_node = static_cast(&expr); + llvm::Value* pred = EmitExpressionImpl(b, *select_node->get_pred()); + llvm::Value* v_true = + EmitExpressionImpl(b, *select_node->get_on_true()); + llvm::Value* v_false = + EmitExpressionImpl(b, *select_node->get_on_false()); + llvm::Value* nonzero = b->CreateICmpNE( + pred, llvm::ConstantInt::get(i64Type, 0, true), "select_dims_pred"); + return b->CreateSelect(nonzero, v_true, v_false, "select_dims"); + } return nullptr; } diff --git a/third_party/xla/xla/service/shape_inference.cc b/third_party/xla/xla/service/shape_inference.cc index 565093a7d81a62..df8d90b198f6c5 100644 --- a/third_party/xla/xla/service/shape_inference.cc +++ b/third_party/xla/xla/service/shape_inference.cc @@ -228,7 +228,8 @@ absl::StatusOr InferWindowOutputShape(const Shape& base_shape, const int64_t input_dimension = ShapeUtil::GetDimension(base_shape, i); const DExpr& input_expression = base_shape.expressions(i); - + const int64_t dilated_window = + window_util::DilatedBound(dim.size(), dim.window_dilation()); if (IsUnboundedDynamicSize(input_dimension)) { output_dimensions[i] = Shape::kUnboundedSize; } else { @@ -236,9 +237,6 @@ absl::StatusOr InferWindowOutputShape(const Shape& base_shape, input_dimension, dim.base_dilation()); const int64_t padded_dilated_base = dim.padding_low() + dilated_base + dim.padding_high(); - const int64_t dilated_window = - window_util::DilatedBound(dim.size(), dim.window_dilation()); - output_dimensions[i] = window_util::StridedBound( padded_dilated_base, dilated_window, dim.stride()); } @@ -248,15 +246,20 @@ absl::StatusOr InferWindowOutputShape(const Shape& base_shape, continue; } - DExpr dilated_base_expr = - ((dim.base_dilation() * (input_expression - 1)) + 1).simplify(); + DExpr dilated_base_expr = input_expression; + if (dim.base_dilation() != 1) { + dilated_base_expr = + DExpr::Max((dim.base_dilation() * (input_expression - 1)) + 1, + DExpr::Const(0)) + .simplify(); + } DExpr padded_dilated_base_expr = (dilated_base_expr + dim.padding_low() + dim.padding_high()).simplify(); - DExpr dilated_window_expr = - DExpr::Const(window_util::DilatedBound(dim.size(), dim.window_dilation())); + DExpr strided_bound_expr = + (padded_dilated_base_expr - dilated_window + 1 + dim.stride() - 1) / + dim.stride(); output_expressions[i] = - (((padded_dilated_base_expr - dilated_window_expr) / dim.stride()) + 1) - .simplify(); + DExpr::Max(strided_bound_expr, DExpr::Const(0)).simplify(); } return ShapeUtil::MakeValidatedShape(element_type, output_dimensions, diff --git a/third_party/xla/xla/service/shape_inference_test.cc b/third_party/xla/xla/service/shape_inference_test.cc index 7efbfbb4684f5a..b4ed0eaa1cdfa6 100644 --- a/third_party/xla/xla/service/shape_inference_test.cc +++ b/third_party/xla/xla/service/shape_inference_test.cc @@ -1518,6 +1518,58 @@ TEST_F(ReduceShapeInferenceTest, ReduceWindowMultiOutput) { *inferred_shape)); } +TEST_F(ReduceShapeInferenceTest, + ReduceWindowPreservesDynamicStridedBoundExpression) { + Shape operand = ShapeUtil::MakeShape( + F32, {101}, std::vector{true}, {DExpr::Var(1)}); + Window window; + WindowDimension* dimension = window.add_dimensions(); + dimension->set_size(3); + dimension->set_stride(2); + dimension->set_padding_low(1); + dimension->set_padding_high(1); + dimension->set_base_dilation(1); + dimension->set_window_dilation(1); + + TF_ASSERT_OK_AND_ASSIGN( + Shape inferred, + ShapeInference::InferReduceWindowShape(operand, f32_, window)); + EXPECT_EQ(51, inferred.dimensions(0)); + EXPECT_TRUE(inferred.expressions(0) == + DExpr::Max((DExpr::Var(1) + 1) / 2, DExpr::Const(0))); + + DExpr runtime_expression = + inferred.expressions(0).substitute(1, DExpr::Const(100)).simplify(); + ASSERT_TRUE(runtime_expression->is_constant()); + EXPECT_EQ(50, runtime_expression->get_val()); +} + +TEST_F(ReduceShapeInferenceTest, + ReduceWindowClampsDynamicStridedBoundAtZero) { + Shape operand = ShapeUtil::MakeShape( + F32, {101}, std::vector{true}, {DExpr::Var(2)}); + Window window; + WindowDimension* dimension = window.add_dimensions(); + dimension->set_size(5); + dimension->set_stride(1); + dimension->set_padding_low(0); + dimension->set_padding_high(0); + dimension->set_base_dilation(1); + dimension->set_window_dilation(1); + + TF_ASSERT_OK_AND_ASSIGN( + Shape inferred, + ShapeInference::InferReduceWindowShape(operand, f32_, window)); + EXPECT_EQ(97, inferred.dimensions(0)); + EXPECT_TRUE(inferred.expressions(0) == + DExpr::Max(DExpr::Var(2) - 4, DExpr::Const(0))); + + DExpr runtime_expression = + inferred.expressions(0).substitute(2, DExpr::Const(2)).simplify(); + ASSERT_TRUE(runtime_expression->is_constant()); + EXPECT_EQ(0, runtime_expression->get_val()); +} + TEST_F(ReduceShapeInferenceTest, ErrorMultiOutputBadReducerInput1) { const Shape f32_arg_shape = ShapeUtil::MakeShape(F32, {5, 3}); const Shape s32_arg_shape = ShapeUtil::MakeShape(S32, {5, 3}); diff --git a/third_party/xla/xla/shape_expr.cc b/third_party/xla/xla/shape_expr.cc index 5beab5cedbca75..ed40b6af5ec7bb 100644 --- a/third_party/xla/xla/shape_expr.cc +++ b/third_party/xla/xla/shape_expr.cc @@ -208,6 +208,12 @@ std::optional ToCanonicalAffine(const DynExpr* expr) { } return MultiplyAffineByRational(*lhs, rhs->denominator, rhs->constant); } + case DExpr::Kind::kMax: + case DExpr::Kind::kGt: + case DExpr::Kind::kSelect: + return std::nullopt; + default: + return std::nullopt; } return std::nullopt; } @@ -345,8 +351,56 @@ std::unique_ptr SimplifyFallback(const DynExpr* expr) { } return std::make_unique
(lhs.release(), rhs.release()); } + case DExpr::Kind::kMax: { + const auto* max = static_cast(expr); + auto lhs = std::unique_ptr(max->get_lhs()->s()); + auto rhs = std::unique_ptr(max->get_rhs()->s()); + if (lhs->kind() == DExpr::Kind::kUnknown || + rhs->kind() == DExpr::Kind::kUnknown) { + return std::make_unique(); + } + Constant* l = AsConstant(lhs.get()); + Constant* r = AsConstant(rhs.get()); + if (l && r) { + return std::make_unique( + std::max(l->get_val(), r->get_val())); + } + if (l && l->get_val() == 0 && + rhs->kind() == DExpr::Kind::kVariable) { + return rhs; + } + if (r && r->get_val() == 0 && + lhs->kind() == DExpr::Kind::kVariable) { + return lhs; + } + if (*lhs == *rhs) return lhs; + return std::make_unique(lhs.release(), rhs.release()); + } + case DExpr::Kind::kGt: { + const auto* gt = static_cast(expr); + auto lhs = std::unique_ptr(gt->get_lhs()->s()); + auto rhs = std::unique_ptr(gt->get_rhs()->s()); + if (lhs->is_constant() && rhs->is_constant()) { + return std::make_unique(lhs->get_val() > rhs->get_val()); + } + if (*lhs == *rhs) return std::make_unique(0); + return std::make_unique(lhs.release(), rhs.release()); + } + case DExpr::Kind::kSelect: { + const auto* select = static_cast(expr); + auto pred = std::unique_ptr(select->get_pred()->s()); + auto on_true = std::unique_ptr(select->get_on_true()->s()); + auto on_false = std::unique_ptr(select->get_on_false()->s()); + if (*on_true == *on_false) return on_true; + if (pred->is_constant()) { + return pred->get_val() != 0 ? std::move(on_true) : std::move(on_false); + } + return std::make_unique(pred.release(), on_true.release(), + on_false.release()); + } + default: + return expr->clone(); } - return expr->clone(); } std::unique_ptr SimplifyCanonical(const DynExpr* expr) { @@ -401,6 +455,21 @@ bool operator<(DynExpr& lhs, int64_t d) { return lhs.is_constant() && lhs.get_val() < d; } +DExpr DExpr::Max(const DExpr& lhs, const DExpr& rhs) { + return Adopt(new xla::MaxExpr(lhs.clone().release(), rhs.clone().release())); +} + +DExpr DExpr::Gt(const DExpr& lhs, const DExpr& rhs) { + return Adopt(new xla::GtExpr(lhs.clone().release(), rhs.clone().release())); +} + +DExpr DExpr::Select(const DExpr& pred, const DExpr& on_true, + const DExpr& on_false) { + return Adopt(new xla::SelectExpr(pred.clone().release(), + on_true.clone().release(), + on_false.clone().release())); +} + bool DynExpr::equal(DynExpr* expr1, DynExpr* expr2) { auto e1 = std::unique_ptr(expr1->s()); auto e2 = std::unique_ptr(expr2->s()); @@ -464,6 +533,29 @@ bool DynExpr::equal(DynExpr* expr1, DynExpr* expr2) { auto* d = cd->get_rhs(); return *a == *c && *b == *d; } + if (e1->kind() == DExpr::Kind::kMax && e2->kind() == DExpr::Kind::kMax) { + auto* ab = static_cast(e1.get()); + auto* cd = static_cast(e2.get()); + auto* a = ab->get_lhs(); + auto* b = ab->get_rhs(); + auto* c = cd->get_lhs(); + auto* d = cd->get_rhs(); + return (*a == *c && *b == *d) || (*a == *d && *b == *c); + } + if (e1->kind() == DExpr::Kind::kGt && e2->kind() == DExpr::Kind::kGt) { + auto* lhs = static_cast(e1.get()); + auto* rhs = static_cast(e2.get()); + return *lhs->get_lhs() == *rhs->get_lhs() && + *lhs->get_rhs() == *rhs->get_rhs(); + } + if (e1->kind() == DExpr::Kind::kSelect && + e2->kind() == DExpr::Kind::kSelect) { + auto* lhs = static_cast(e1.get()); + auto* rhs = static_cast(e2.get()); + return *lhs->get_pred() == *rhs->get_pred() && + *lhs->get_on_true() == *rhs->get_on_true() && + *lhs->get_on_false() == *rhs->get_on_false(); + } return false; } @@ -479,6 +571,12 @@ DynExpr* Sub::s() { return SimplifyCanonical(this).release(); } DynExpr* Div::s() { return SimplifyCanonical(this).release(); } +DynExpr* MaxExpr::s() { return SimplifyCanonical(this).release(); } + +DynExpr* GtExpr::s() { return SimplifyCanonical(this).release(); } + +DynExpr* SelectExpr::s() { return SimplifyCanonical(this).release(); } + std::ostream& operator<<(std::ostream& os, DynExpr* expr) { auto simplified = std::unique_ptr(expr->s()); StringPrinter printer; diff --git a/third_party/xla/xla/shape_expr.h b/third_party/xla/xla/shape_expr.h index feecb054bfbc92..b503ff70690dcc 100644 --- a/third_party/xla/xla/shape_expr.h +++ b/third_party/xla/xla/shape_expr.h @@ -16,15 +16,18 @@ limitations under the License. #ifndef XLA_SHAPE_EXPR_H_ #define XLA_SHAPE_EXPR_H_ +#include #include #include #include #include #include +#include #include #include "absl/hash/hash.h" #include "absl/log/check.h" +#include "absl/log/log.h" #include "absl/types/span.h" #include "xla/printer.h" #include "xla/xla_data.pb.h" @@ -44,6 +47,9 @@ enum class DExprKind { kSub, kMul, kDiv, + kMax, + kGt, + kSelect, }; class DynExpr { @@ -100,6 +106,10 @@ class DExpr { static DExpr Adopt(DynExpr* expr) { return DExpr(std::unique_ptr(expr)); } static DExpr Const(int64_t value) { return Adopt(DynExpr::_(value)); } static DExpr Var(int var_id) { return Adopt(DynExpr::V(var_id)); } + static DExpr Max(const DExpr& lhs, const DExpr& rhs); + static DExpr Gt(const DExpr& lhs, const DExpr& rhs); + static DExpr Select(const DExpr& pred, const DExpr& on_true, + const DExpr& on_false); bool is_unknown() const { return expr_ != nullptr && expr_->kind() == DExprKind::kUnknown; } @@ -185,8 +195,7 @@ class UnknownExpr : public DynExpr { return clone().release(); } std::set get_all_ids() override { return {}; } - std::optional solve(int64_t x) override { - (void)x; + std::optional solve(int64_t) override { return std::nullopt; } DynExpr* s() override { return clone().release(); } @@ -529,6 +538,163 @@ class Div : public DynExpr { ~Div() override = default; }; +// max(lhs, rhs) +class MaxExpr : public DynExpr { + std::unique_ptr lhs; + std::unique_ptr rhs; + + public: + MaxExpr(DynExpr* l, DynExpr* r) : lhs(l), rhs(r) {} + std::unique_ptr clone() const override { + return std::make_unique(lhs->clone().release(), + rhs->clone().release()); + } + DExprKind kind() const override { return DExprKind::kMax; } + void print(xla::Printer* printer) const override { + printer->Append("max("); + lhs->print(printer); + printer->Append(", "); + rhs->print(printer); + printer->Append(")"); + } + void to_proto(xla::ExpressionProto* proto) const override { + auto* max_msg = proto->mutable_max_node(); + lhs->to_proto(max_msg->mutable_lhs()); + rhs->to_proto(max_msg->mutable_rhs()); + } + bool is_constant() const override { + return lhs->is_constant() && rhs->is_constant(); + } + int64_t get_val() const override { + return std::max(lhs->get_val(), rhs->get_val()); + } + DynExpr* get_lhs() const { return lhs.get(); } + DynExpr* get_rhs() const { return rhs.get(); } + DynExpr* substitute(int id, DynExpr* v) override { + return new MaxExpr(lhs->substitute(id, v), rhs->substitute(id, v)); + } + std::set get_all_ids() override { + auto ids = lhs->get_all_ids(); + ids.merge(rhs->get_all_ids()); + return ids; + } + // Max is not invertible: either operand may have produced the result. + std::optional solve(int64_t x) override { + StringPrinter printer; + print(&printer); + LOG(WARNING) << "Cannot solve Max dynamic shape expression for value " << x + << ": " << std::move(printer).ToString(); + return std::nullopt; + } + DynExpr* s() override; +}; + +class GtExpr : public DynExpr { + std::unique_ptr lhs; + std::unique_ptr rhs; + + public: + GtExpr(DynExpr* l, DynExpr* r) : lhs(l), rhs(r) {} + std::unique_ptr clone() const override { + return std::make_unique(lhs->clone().release(), + rhs->clone().release()); + } + DExprKind kind() const override { return DExprKind::kGt; } + void print(xla::Printer* printer) const override { + printer->Append("("); + lhs->print(printer); + printer->Append(" > "); + rhs->print(printer); + printer->Append(")"); + } + void to_proto(xla::ExpressionProto* proto) const override { + auto* gt_msg = proto->mutable_gt_node(); + lhs->to_proto(gt_msg->mutable_lhs()); + rhs->to_proto(gt_msg->mutable_rhs()); + } + bool is_constant() const override { + return lhs->is_constant() && rhs->is_constant(); + } + int64_t get_val() const override { return lhs->get_val() > rhs->get_val(); } + DynExpr* get_lhs() const { return lhs.get(); } + DynExpr* get_rhs() const { return rhs.get(); } + DynExpr* substitute(int id, DynExpr* v) override { + return new GtExpr(lhs->substitute(id, v), rhs->substitute(id, v)); + } + std::set get_all_ids() override { + auto ids = lhs->get_all_ids(); + ids.merge(rhs->get_all_ids()); + return ids; + } + std::optional solve(int64_t x) override { + StringPrinter printer; + print(&printer); + LOG(WARNING) << "Cannot solve Gt dynamic shape expression for value " << x + << ": " << std::move(printer).ToString(); + return std::nullopt; + } + DynExpr* s() override; +}; + +class SelectExpr : public DynExpr { + std::unique_ptr pred; + std::unique_ptr on_true; + std::unique_ptr on_false; + + public: + SelectExpr(DynExpr* p, DynExpr* t, DynExpr* f) + : pred(p), on_true(t), on_false(f) {} + std::unique_ptr clone() const override { + return std::make_unique(pred->clone().release(), + on_true->clone().release(), + on_false->clone().release()); + } + DExprKind kind() const override { return DExprKind::kSelect; } + void print(xla::Printer* printer) const override { + printer->Append("select("); + pred->print(printer); + printer->Append(", "); + on_true->print(printer); + printer->Append(", "); + on_false->print(printer); + printer->Append(")"); + } + void to_proto(xla::ExpressionProto* proto) const override { + auto* select_msg = proto->mutable_select_node(); + pred->to_proto(select_msg->mutable_pred()); + on_true->to_proto(select_msg->mutable_on_true()); + on_false->to_proto(select_msg->mutable_on_false()); + } + bool is_constant() const override { + return pred->is_constant() && on_true->is_constant() && + on_false->is_constant(); + } + int64_t get_val() const override { + return pred->get_val() != 0 ? on_true->get_val() : on_false->get_val(); + } + DynExpr* get_pred() const { return pred.get(); } + DynExpr* get_on_true() const { return on_true.get(); } + DynExpr* get_on_false() const { return on_false.get(); } + DynExpr* substitute(int id, DynExpr* v) override { + return new SelectExpr(pred->substitute(id, v), on_true->substitute(id, v), + on_false->substitute(id, v)); + } + std::set get_all_ids() override { + auto ids = pred->get_all_ids(); + ids.merge(on_true->get_all_ids()); + ids.merge(on_false->get_all_ids()); + return ids; + } + std::optional solve(int64_t x) override { + StringPrinter printer; + print(&printer); + LOG(WARNING) << "Cannot solve Select dynamic shape expression for value " + << x << ": " << std::move(printer).ToString(); + return std::nullopt; + } + DynExpr* s() override; +}; + DynExpr* operator*(DynExpr& lhs, DynExpr& rhs); DynExpr* operator*(int64_t k, DynExpr& rhs); DynExpr* operator/(DynExpr& lhs, DynExpr& rhs); @@ -593,6 +759,21 @@ inline DExpr DExprFromProto(const xla::ExpressionProto& proto) { const auto& div = proto.div_node(); return DExprFromProto(div.lhs()) / DExprFromProto(div.rhs()); } + case ExpressionProto::kMaxNode: { + const auto& max = proto.max_node(); + return DExpr::Max(DExprFromProto(max.lhs()), + DExprFromProto(max.rhs())); + } + case ExpressionProto::kGtNode: { + const auto& gt = proto.gt_node(); + return DExpr::Gt(DExprFromProto(gt.lhs()), DExprFromProto(gt.rhs())); + } + case ExpressionProto::kSelectNode: { + const auto& select = proto.select_node(); + return DExpr::Select(DExprFromProto(select.pred()), + DExprFromProto(select.on_true()), + DExprFromProto(select.on_false())); + } case ExpressionProto::NODE_TYPE_NOT_SET: default: return DExpr::Unknown(kMissingExpressionSentinel); diff --git a/third_party/xla/xla/shape_test.cc b/third_party/xla/xla/shape_test.cc index b6e4bcd79c81bb..923f106adba26c 100644 --- a/third_party/xla/xla/shape_test.cc +++ b/third_party/xla/xla/shape_test.cc @@ -53,9 +53,10 @@ class ShapeTest : public ::testing::Test { const Shape nested_tuple_ = ShapeUtil::MakeTupleShape({tuple_, matrix_, token_}); const Shape dynamic_matrix_ = - ShapeUtil::MakeShape(S32, {5, 2}, {true, false}); + ShapeUtil::MakeShape(S32, {5, 2}, std::vector{true, false}, {}); const Shape unbounded_ = - ShapeUtil::MakeShape(F32, {Shape::kUnboundedSize, 784}, {true, false}); + ShapeUtil::MakeShape(F32, {Shape::kUnboundedSize, 784}, + std::vector{true, false}, {}); }; // Tests that if the dynamic_dimensions parameter empty in the Shape @@ -105,8 +106,8 @@ TEST_F(ShapeTest, ShapeToString) { } TEST_F(ShapeTest, DynamicShapeToString) { - Shape array_shape = - ShapeUtil::MakeShape(F32, {23, 44, 55}, {true, false, true}); + Shape array_shape = ShapeUtil::MakeShape( + F32, {23, 44, 55}, std::vector{true, false, true}, {}); EXPECT_EQ("f32[<=23,44,<=55]", array_shape.ToString()); array_shape.set_dynamic_dimension(2, false); @@ -125,6 +126,36 @@ TEST_F(ShapeTest, DExprSimplifyCombinesEqualFractions) { EXPECT_EQ("A", DExprToString(expr.simplify())); } +TEST_F(ShapeTest, DExprMaxSimplifiesAndRoundTrips) { + DExpr expr = DExpr::Max(DExpr::Var(1), DExpr::Const(4)); + EXPECT_EQ("max(A, 4)", DExprToString(expr.simplify())); + EXPECT_FALSE(expr->solve(7).has_value()); + + DExpr clamped = DExpr::Max(DExpr::Var(1), DExpr::Const(0)); + EXPECT_EQ("A", DExprToString(clamped.simplify())); + + DExpr evaluated = expr.substitute(1, DExpr::Const(7)).simplify(); + EXPECT_EQ(DExpr::Kind::kConstant, evaluated.kind()); + EXPECT_EQ(7, evaluated->get_val()); + + ExpressionProto proto; + expr.to_proto(&proto); + EXPECT_TRUE(expr == DExprFromProto(proto)); +} + +TEST_F(ShapeTest, DExprSelectUsesDynamicPredicate) { + DExpr delta = DExpr::Var(1); + DExpr expr = DExpr::Select(DExpr::Gt(delta, DExpr::Const(0)), + DExpr::Const(7), DExpr::Const(11)); + EXPECT_EQ("select((A > 0), 7, 11)", DExprToString(expr.simplify())); + EXPECT_EQ(7, expr.substitute(1, DExpr::Const(2))->s()->get_val()); + EXPECT_EQ(11, expr.substitute(1, DExpr::Const(-2))->s()->get_val()); + + ExpressionProto proto; + expr.to_proto(&proto); + EXPECT_TRUE(expr == DExprFromProto(proto)); +} + TEST_F(ShapeTest, DeleteDimensions) { Shape shape = ShapeUtil::MakeShapeWithDenseLayout(F32, {5, 3, 2, 7, 9}, {2, 0, 1, 4, 3}); diff --git a/third_party/xla/xla/xla_data.proto b/third_party/xla/xla/xla_data.proto index 798dc12fb5734e..169e82bef57aed 100644 --- a/third_party/xla/xla/xla_data.proto +++ b/third_party/xla/xla/xla_data.proto @@ -1215,6 +1215,9 @@ message ExpressionProto { SubNode sub_node = 4; // exp - exp MulNode mul_node = 5; // exp * exp DivNode div_node = 6; // exp / exp + MaxNode max_node = 7; // max(exp, exp) + GtNode gt_node = 8; // exp > exp + SelectNode select_node = 9; // select(pred, on_true, on_false) } } @@ -1236,4 +1239,20 @@ message MulNode { message DivNode { ExpressionProto lhs = 1; ExpressionProto rhs = 2; -} \ No newline at end of file +} + +message MaxNode { + ExpressionProto lhs = 1; + ExpressionProto rhs = 2; +} + +message GtNode { + ExpressionProto lhs = 1; + ExpressionProto rhs = 2; +} + +message SelectNode { + ExpressionProto pred = 1; + ExpressionProto on_true = 2; + ExpressionProto on_false = 3; +} From 7f3d003f381ae22371ba59b8e83cce115b5f0a6c Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Wed, 29 Jul 2026 13:03:15 +0100 Subject: [PATCH 03/15] Preserve dynamic shape expressions across XLA shape reconstruction (#13) --- .../compiler/tf2xla/kernels/resampler_ops.cc | 17 ++++-- .../collectives/collective_quantizer.cc | 1 + .../transforms/expanders/reduce_decomposer.cc | 1 + .../service/gpu/transforms/gemm_rewriter.cc | 2 + .../gpu/transforms/windowed_einsum_handler.cc | 2 + .../xla/xla/service/layout_assignment.cc | 2 + third_party/xla/xla/shape_util.cc | 54 ++++++++++++++----- third_party/xla/xla/shape_util.h | 11 ++++ third_party/xla/xla/shape_util_test.cc | 27 ++++++++++ 9 files changed, 99 insertions(+), 18 deletions(-) diff --git a/tensorflow/compiler/tf2xla/kernels/resampler_ops.cc b/tensorflow/compiler/tf2xla/kernels/resampler_ops.cc index c54c4613d29e44..1d2f23efa62d3a 100644 --- a/tensorflow/compiler/tf2xla/kernels/resampler_ops.cc +++ b/tensorflow/compiler/tf2xla/kernels/resampler_ops.cc @@ -122,12 +122,15 @@ XlaOp ConcatenateIota(xla::XlaBuilder* b, XlaOp indices, for (auto dim : warp_shape) { dimensions.push_back(dim.size); } + std::vector expressions(warp_shape.get_expressions().begin(), + warp_shape.get_expressions().end()); // Except the last dimension, which is of size 1. dimensions.back() = 1; + expressions.back() = xla::DExpr::Const(1); - auto batch_indices = - xla::Iota(b, xla::ShapeUtil::MakeShape(xla::S32, dimensions), - /*iota_dimension=*/0); + auto batch_indices = xla::Iota( + b, xla::ShapeUtil::MakeShape(xla::S32, dimensions, expressions), + /*iota_dimension=*/0); return xla::ConcatInDim(b, {batch_indices, indices}, dimensions.size() - 1); } @@ -365,14 +368,18 @@ XlaOp CalculateGradWarp(XlaOpKernelContext* ctx, XlaOp grad_output, XlaOp ratio, auto warp_dims = warp_shape.dim_sizes(); std::vector warp_dims_without_last_dims(warp_dims.begin(), warp_dims.end() - 1); + std::vector warp_expressions( + warp_shape.get_expressions().begin(), warp_shape.get_expressions().end()); + warp_expressions.pop_back(); // With dimension [batch, dim_0, ...dim_n, 4] std::vector neighbor_broadcast_dims = warp_dims_without_last_dims; neighbor_broadcast_dims.push_back(4); + warp_expressions.push_back(xla::DExpr::Const(4)); // With dimension [batch, dim_0, ...dim_n, 4] - auto neighbor_broadcast_shape = - xla::ShapeUtil::MakeShape(data_type, neighbor_broadcast_dims); + auto neighbor_broadcast_shape = xla::ShapeUtil::MakeShape( + data_type, neighbor_broadcast_dims, warp_expressions); const int64_t last_warp_dim = warp_shape.dims() - 1; diff --git a/third_party/xla/xla/hlo/transforms/collectives/collective_quantizer.cc b/third_party/xla/xla/hlo/transforms/collectives/collective_quantizer.cc index b3c2ffe79ec00c..e038b7338a1bc2 100644 --- a/third_party/xla/xla/hlo/transforms/collectives/collective_quantizer.cc +++ b/third_party/xla/xla/hlo/transforms/collectives/collective_quantizer.cc @@ -121,6 +121,7 @@ HloInstruction* ApplyUnaries(HloInstruction* instr, instr = instr->AddInstruction(unary->CloneWithNewOperands( ShapeUtil::MakeShapeWithDenseLayout( instr->shape().element_type(), unary->shape().dimensions(), + unary->shape().expressions(), unary->shape().layout().minor_to_major()), {instr})); } diff --git a/third_party/xla/xla/hlo/transforms/expanders/reduce_decomposer.cc b/third_party/xla/xla/hlo/transforms/expanders/reduce_decomposer.cc index 2fe502429287b4..ca7eebb9cfb7e8 100644 --- a/third_party/xla/xla/hlo/transforms/expanders/reduce_decomposer.cc +++ b/third_party/xla/xla/hlo/transforms/expanders/reduce_decomposer.cc @@ -47,6 +47,7 @@ class VariadicReductionLayoutEqualizer : public DfsHloRewriteVisitor { if (first_input_s.layout() != input_s.layout()) { Shape new_input_s = ShapeUtil::MakeShapeWithDenseLayout( input_s.element_type(), input_s.dimensions(), + input_s.expressions(), first_input_s.layout().minor_to_major()); auto copy = MakeCopyHlo(input, new_input_s); changed = true; diff --git a/third_party/xla/xla/service/gpu/transforms/gemm_rewriter.cc b/third_party/xla/xla/service/gpu/transforms/gemm_rewriter.cc index 5054a440778105..484a1e75ac0aeb 100644 --- a/third_party/xla/xla/service/gpu/transforms/gemm_rewriter.cc +++ b/third_party/xla/xla/service/gpu/transforms/gemm_rewriter.cc @@ -1316,6 +1316,7 @@ class GemmRewriterVisitor : public DfsHloRewriteVisitor { x = instr->AddInstruction(op.first->CloneWithNewOperands( ShapeUtil::MakeShapeWithDenseLayout( x->shape().element_type(), op.first->shape().dimensions(), + op.first->shape().expressions(), op.first->shape().layout().minor_to_major()), operands)); } @@ -1378,6 +1379,7 @@ class GemmRewriterVisitor : public DfsHloRewriteVisitor { instr->AddInstruction(HloInstruction::CreateCustomCall( ShapeUtil::MakeShapeWithDenseLayout( instr->shape().element_type(), new_output_shape.dimensions(), + new_output_shape.expressions(), instr->shape().layout().minor_to_major()), operands_list, kCublasLtMatmulF8CallTarget)); TF_RETURN_IF_ERROR(new_custom_call->set_backend_config(gpu_backend_config)); diff --git a/third_party/xla/xla/service/gpu/transforms/windowed_einsum_handler.cc b/third_party/xla/xla/service/gpu/transforms/windowed_einsum_handler.cc index ce454624144803..4585af7203946b 100644 --- a/third_party/xla/xla/service/gpu/transforms/windowed_einsum_handler.cc +++ b/third_party/xla/xla/service/gpu/transforms/windowed_einsum_handler.cc @@ -183,11 +183,13 @@ absl::StatusOr ShiftDequantizationF8( for (HloInstruction* unary : unaries[k]) { Shape new_shape = ShapeUtil::MakeShapeWithDenseLayout( operands[k]->shape().element_type(), unary->shape().dimensions(), + unary->shape().expressions(), unary->shape().layout().minor_to_major()); operands[k] = unary->AddInstruction(unary->CloneWithNewOperands( ShapeUtil::MakeShapeWithDenseLayout( operands[k]->shape().element_type(), unary->shape().dimensions(), + unary->shape().expressions(), unary->shape().layout().minor_to_major()), {operands[k]})); } diff --git a/third_party/xla/xla/service/layout_assignment.cc b/third_party/xla/xla/service/layout_assignment.cc index b5adc212b53a17..001a873118ab7f 100644 --- a/third_party/xla/xla/service/layout_assignment.cc +++ b/third_party/xla/xla/service/layout_assignment.cc @@ -1401,6 +1401,7 @@ std::unique_ptr LayoutAssignment::ChooseOperandLayoutFromOutputLayout( const Shape& output_shape = instruction->shape(); Shape output_shape_with_layout = ShapeUtil::MakeShapeWithDenseLayout( output_shape.element_type(), output_shape.dimensions(), + output_shape.expressions(), LayoutUtil::MinorToMajor(output_layout)); Shape operand_shape = operand->shape(); *operand_shape.mutable_layout() = @@ -1539,6 +1540,7 @@ std::unique_ptr LayoutAssignment::ChooseOutputLayoutFromOperandLayout( } Shape operand_shape_with_layout = ShapeUtil::MakeShapeWithDenseLayout( operand->shape().element_type(), operand->shape().dimensions(), + operand->shape().expressions(), LayoutUtil::MinorToMajor(operand_layout)); Shape output_shape = user->shape(); *output_shape.mutable_layout() = diff --git a/third_party/xla/xla/shape_util.cc b/third_party/xla/xla/shape_util.cc index 5f8b140f873f1b..75970dbeceb325 100644 --- a/third_party/xla/xla/shape_util.cc +++ b/third_party/xla/xla/shape_util.cc @@ -123,6 +123,7 @@ void PrintBufferShape(Printer* printer, const Shape& shape) { // its Layout. absl::StatusOr MakeShapeWithLayoutInternal( PrimitiveType element_type, absl::Span dimensions, + absl::Span expressions, absl::Span minor_to_major, absl::Span tiles, int64_t tail_padding_alignment_in_elements, PrimitiveType index_primitive_type, PrimitiveType pointer_primitive_type, @@ -139,7 +140,8 @@ absl::StatusOr MakeShapeWithLayoutInternal( PrimitiveType_Name(element_type)); } TF_ASSIGN_OR_RETURN(Shape shape, - ShapeUtil::MakeValidatedShape(element_type, dimensions)); + ShapeUtil::MakeValidatedShape(element_type, dimensions, + expressions)); if (element_size_in_bits == ShapeUtil::ByteSizeOfPrimitiveType(element_type) * 8) { // Only set element_size_in_bits if it's different from the default value. @@ -383,11 +385,29 @@ static std::vector MakeExpressions( /* static */ Shape ShapeUtil::MakeShapeWithDenseLayout( PrimitiveType element_type, absl::Span dimensions, + absl::Span expressions, absl::Span minor_to_major, absl::Span tiles, int64_t tail_padding_alignment_in_elements, int64_t element_size_in_bits, int64_t memory_space, absl::Span split_configs) { auto ret = MakeShapeWithLayoutInternal( - element_type, dimensions, minor_to_major, tiles, + element_type, dimensions, expressions, minor_to_major, tiles, + tail_padding_alignment_in_elements, + /*index_primitive_type=*/PRIMITIVE_TYPE_INVALID, + /*pointer_primitive_type=*/PRIMITIVE_TYPE_INVALID, element_size_in_bits, + memory_space, split_configs, + /*physical_shape=*/std::nullopt); + TF_CHECK_OK(ret.status()); + return *ret; +} + +/* static */ Shape ShapeUtil::MakeShapeWithDenseLayout( + PrimitiveType element_type, absl::Span dimensions, + absl::Span minor_to_major, absl::Span tiles, + int64_t tail_padding_alignment_in_elements, int64_t element_size_in_bits, + int64_t memory_space, absl::Span split_configs) { + auto ret = MakeShapeWithLayoutInternal( + element_type, dimensions, MakeExpressions(dimensions), minor_to_major, + tiles, tail_padding_alignment_in_elements, /*index_primitive_type=*/PRIMITIVE_TYPE_INVALID, /*pointer_primitive_type=*/PRIMITIVE_TYPE_INVALID, element_size_in_bits, @@ -404,7 +424,7 @@ static std::vector MakeExpressions( int64_t tail_padding_alignment_in_elements, int64_t element_size_in_bits, int64_t memory_space, std::optional physical_shape) { auto ret = MakeShapeWithLayoutInternal( - element_type, dimensions, minor_to_major, + element_type, dimensions, MakeExpressions(dimensions), minor_to_major, /*tiles=*/{}, tail_padding_alignment_in_elements, index_primitive_type, pointer_primitive_type, element_size_in_bits, memory_space, /*split_configs=*/{}, std::move(physical_shape)); @@ -441,13 +461,9 @@ static std::vector MakeExpressions( /* static */ Shape ShapeUtil::MakeShapeWithDescendingLayout( PrimitiveType element_type, absl::Span dimensions, absl::Span expressions) { - auto shape = MakeShapeWithDenseLayout(element_type, dimensions, - LayoutUtil::MakeDescendingLayout( - dimensions.size()) - .minor_to_major()); - std::vector exprs(expressions.begin(), expressions.end()); - shape.set_expressions(exprs); - return shape; + return MakeShapeWithDenseLayout( + element_type, dimensions, expressions, + LayoutUtil::MakeDescendingLayout(dimensions.size()).minor_to_major()); } /* static */ Shape @@ -461,7 +477,17 @@ ShapeUtil::MakeShapeWithDescendingLayoutAndSamePhysicalLayout( } dims[i] = shape.dimensions(dim); } - Shape new_shape = MakeShapeWithDescendingLayout(shape.element_type(), dims); + std::vector expressions; + expressions.reserve(shape.dimensions().size()); + for (int i = 0; i < shape.dimensions().size(); ++i) { + int dim = i; + if (shape.has_layout()) { + dim = LayoutUtil::Major(shape.layout(), dim); + } + expressions.push_back(shape.expressions(dim)); + } + Shape new_shape = MakeShapeWithDescendingLayout(shape.element_type(), dims, + expressions); // Since the physical layout is kept the same, the tiles and element size are // the same also. if (shape.has_layout()) { @@ -1839,7 +1865,8 @@ ShapeUtil::DecomposeBitcastToTrt(const Shape& input_shape, } Shape output_shape_with_layout = MakeShapeWithDenseLayout( - output_shape.element_type(), output_shape.dimensions(), output_layout); + output_shape.element_type(), output_shape.dimensions(), + output_shape.expressions(), output_layout); CHECK(ReshapeIsBitcast(input_shape, output_shape_with_layout)) << "reshape is not a bitcast for input_shape: " << ShapeUtil::HumanStringWithLayout(input_shape) @@ -2072,7 +2099,8 @@ struct ParallelState { } // Create the shape of the "work" which has same layout as the original shape. - Shape work_shape = ShapeUtil::MakeShape(shape.element_type(), work_dims); + Shape work_shape = ShapeUtil::MakeShape(shape.element_type(), work_dims, + shape.expressions()); *work_shape.mutable_layout() = shape.layout(); // We target one task (partition) per available thread. diff --git a/third_party/xla/xla/shape_util.h b/third_party/xla/xla/shape_util.h index c84678fa88c9f9..d7b797ef97fd3b 100644 --- a/third_party/xla/xla/shape_util.h +++ b/third_party/xla/xla/shape_util.h @@ -453,6 +453,17 @@ class ShapeUtil { dimensions); } + // Constructs a new dense array shape with the given minor_to_major order in + // its Layout. Returns a value shape such that shape.has_layout(). + static Shape MakeShapeWithDenseLayout( + PrimitiveType element_type, absl::Span dimensions, + absl::Span expressions, + absl::Span minor_to_major, + absl::Span tiles = {}, + int64_t tail_padding_alignment_in_elements = 1, + int64_t element_size_in_bits = 0, int64_t memory_space = 0, + absl::Span split_configs = {}); + // Constructs a new dense array shape with the given minor_to_major order in // its Layout. Returns a value shape such that shape.has_layout(). static Shape MakeShapeWithDenseLayout( diff --git a/third_party/xla/xla/shape_util_test.cc b/third_party/xla/xla/shape_util_test.cc index 743aa8ecf4d507..64316af386ba80 100644 --- a/third_party/xla/xla/shape_util_test.cc +++ b/third_party/xla/xla/shape_util_test.cc @@ -1129,6 +1129,15 @@ TEST(ShapeUtilTest, DeleteDimensions) { ShapeUtil::MakeShapeWithDenseLayout(F32, {5, 2}, {1, 0})); } +TEST(ShapeUtilTest, MakeShapeWithDenseLayoutPreservesExpressions) { + std::vector expressions = {DExpr::Var(7), DExpr::Const(24)}; + Shape shape = ShapeUtil::MakeShapeWithDenseLayout( + F32, {10, 24}, expressions, {1, 0}); + + EXPECT_TRUE(shape.expressions(0) == DExpr::Var(7)); + EXPECT_TRUE(shape.expressions(1) == DExpr::Const(24)); +} + TEST(ShapeUtilTest, MakeShapeWithDescendingLayoutAndSamePhysicalLayout) { Shape shape = ShapeUtil::MakeShapeWithDenseLayout(F32, {128, 24, 4, 48, 48}, {2, 4, 3, 1, 0}); @@ -1153,6 +1162,24 @@ TEST(ShapeUtilTest, EXPECT_EQ(new_shape, expected_shape); } +TEST(ShapeUtilTest, + MakeShapeWithDescendingLayoutAndSamePhysicalLayoutPreservesExpressions) { + std::vector expressions = {DExpr::Var(1), DExpr::Const(24), + DExpr::Var(2), DExpr::Const(48), + DExpr::Const(48)}; + Shape shape = ShapeUtil::MakeShapeWithDenseLayout( + F32, {128, 24, 4, 48, 48}, expressions, {2, 4, 3, 1, 0}); + + Shape new_shape = + ShapeUtil::MakeShapeWithDescendingLayoutAndSamePhysicalLayout(shape); + + EXPECT_TRUE(new_shape.expressions(0) == DExpr::Var(1)); + EXPECT_TRUE(new_shape.expressions(1) == DExpr::Const(24)); + EXPECT_TRUE(new_shape.expressions(2) == DExpr::Const(48)); + EXPECT_TRUE(new_shape.expressions(3) == DExpr::Const(48)); + EXPECT_TRUE(new_shape.expressions(4) == DExpr::Var(2)); +} + TEST(ShapeUtilTest, DeduceTransposeDimensionsForBitcast) { Shape input_shape = ShapeUtil::MakeShapeWithDenseLayout(F32, {5, 3}, {1, 0}); Shape output_shape = ShapeUtil::MakeShapeWithDenseLayout(F32, {3, 5}, {0, 1}); From 61d47251865694fcccc44ec88c453ad32e10886f Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Wed, 29 Jul 2026 13:21:26 +0100 Subject: [PATCH 04/15] Reject dynamic shape expressions in MLIR XLA kernels (#19) --- .../compiler/tf2xla/mlir_xla_op_kernel.cc | 67 +++++++++++++++++++ 1 file changed, 67 insertions(+) diff --git a/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.cc b/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.cc index b1a93508d92896..33e9ed240aeaa7 100644 --- a/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.cc +++ b/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.cc @@ -16,6 +16,7 @@ limitations under the License. #include "tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h" #include +#include #include "absl/status/status.h" #include "absl/strings/str_cat.h" @@ -35,7 +36,9 @@ limitations under the License. #include "tensorflow/core/framework/op_requires.h" #include "tensorflow/core/framework/resource_base.h" #include "tensorflow/core/framework/resource_mgr.h" +#include "tensorflow/core/framework/tensor_shape.h" #include "tensorflow/core/framework/types.pb.h" +#include "tensorflow/core/graph/graph.h" #include "tensorflow/core/platform/errors.h" #include "tensorflow/core/platform/refcount.h" #include "tensorflow/core/platform/status.h" @@ -68,6 +71,68 @@ class MLIRContextResource : public ResourceBase { mlir::MLIRContext mlir_ctx_; }; +bool HasDynamicExpressions(const TensorShape& shape) { + for (const auto& expr : shape.get_expressions()) { + if (expr && expr->is_dynamic()) { + return true; + } + } + return false; +} + +bool HasDynamicExpressions(const xla::Shape& shape) { + for (const auto& expr : shape.expressions()) { + if (expr && expr->is_dynamic()) { + return true; + } + } + return false; +} + +absl::Status RejectDynamicShapeExpressionsInMlirXlaOpKernel( + llvm::ArrayRef args, const Graph& graph) { + for (int i = 0; i < args.size(); ++i) { + const auto& shape = args[i].shape; + const bool has_dynamic_exprs = + std::holds_alternative(shape) + ? HasDynamicExpressions(std::get(shape)) + : HasDynamicExpressions(std::get(shape)); + if (has_dynamic_exprs) { + return errors::Unimplemented( + "MlirXlaOpKernel does not support dynamic shape expressions. " + "Argument ", + i, " carries a dynamic expression."); + } + } + + for (Node* node : graph.nodes()) { + for (const auto& name_attr_pair : node->attrs()) { + const auto& attr_name = name_attr_pair.first; + const auto& attr_value = name_attr_pair.second; + auto maybe_reject_shape = [&](const TensorShapeProto& shape_proto) { + const TensorShape shape(shape_proto); + if (!HasDynamicExpressions(shape)) { + return absl::OkStatus(); + } + return errors::Unimplemented( + "MlirXlaOpKernel does not support dynamic shape expressions. " + "Node '", + node->name(), "' attribute '", attr_name, + "' carries a dynamic expression."); + }; + + if (attr_value.value_case() == AttrValue::kShape) { + TF_RETURN_IF_ERROR(maybe_reject_shape(attr_value.shape())); + } else if (attr_value.value_case() == AttrValue::kList) { + for (const auto& shape_proto : attr_value.list().shape()) { + TF_RETURN_IF_ERROR(maybe_reject_shape(shape_proto)); + } + } + } + } + return absl::OkStatus(); +} + } // namespace absl::Status MlirXlaOpKernel::ContextToXlaArgs( @@ -143,6 +208,8 @@ absl::Status MlirXlaOpKernel::ConstructXlaOp(XlaOpKernelContext* ctx) { // Create a graph that wraps the kernel. TF_ASSIGN_OR_RETURN(auto graph, CreateSingleOpGraph(def(), xla_args, result_dtypes)); + TF_RETURN_IF_ERROR( + RejectDynamicShapeExpressionsInMlirXlaOpKernel(xla_args, *graph)); ResourceMgr* res_manager = ctx->op_kernel_context()->resource_manager(); MLIRContextResource* ctx_res; From 12240691779d5786044602c3f666db342c3b9d86 Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Mon, 27 Apr 2026 11:48:24 +0100 Subject: [PATCH 05/15] Preserve symbolic contents during partial constant folding --- tensorflow/compiler/jit/kernels/xla_ops.cc | 30 +- .../compiler/tf2xla/kernels/const_op.cc | 55 ++- .../core/common_runtime/constant_folding.cc | 402 +++++++++++++++++- 3 files changed, 442 insertions(+), 45 deletions(-) diff --git a/tensorflow/compiler/jit/kernels/xla_ops.cc b/tensorflow/compiler/jit/kernels/xla_ops.cc index 1a84226b2133ab..018e5cb9b196a1 100644 --- a/tensorflow/compiler/jit/kernels/xla_ops.cc +++ b/tensorflow/compiler/jit/kernels/xla_ops.cc @@ -553,30 +553,32 @@ absl::Status CompileToLocalExecutable( return; } + auto inferred_shape_it = attr_map.find("user_inferred_shape"); + auto inferred_contents_it = + attr_map.find("user_inferred_value_contents"); bool has_dynamic = false; auto has_dynamic_it = attr_map.find("has_dynamic"); - if (has_dynamic_it == attr_map.end()) { - return; - } - has_dynamic = has_dynamic_it->second.b(); - if (!has_dynamic) { - return; + if (has_dynamic_it != attr_map.end()) { + has_dynamic = has_dynamic_it->second.b(); } - auto inferred_shape_it = attr_map.find("user_inferred_shape"); - if (inferred_shape_it == attr_map.end()) { - VLOG(1) << "XlaCompileOp saw has_dynamic for const arg " - << arg_index << " node=" << node_name - << " but no user_inferred_shape attr"; + if (inferred_contents_it == attr_map.end() && + (!has_dynamic || inferred_shape_it == attr_map.end())) { return; } TensorShapeProto inferred_shape_proto; - inferred_shape_proto = inferred_shape_it->second.shape(); + if (inferred_contents_it != attr_map.end()) { + inferred_shape_proto = inferred_contents_it->second.shape(); + } else { + inferred_shape_proto = inferred_shape_it->second.shape(); + } TensorShape inferred_shape(inferred_shape_proto); - if (!TensorShapeUtils::IsVector(arg.constant_value.shape()) || - arg.constant_value.NumElements() != inferred_shape.dims()) { + if (!((TensorShapeUtils::IsVector(arg.constant_value.shape()) && + arg.constant_value.NumElements() == inferred_shape.dims()) || + (TensorShapeUtils::IsScalar(arg.constant_value.shape()) && + inferred_shape.dims() == 1))) { VLOG(1) << "XlaCompileOp const arg " << arg_index << " node=" << node_name << " has dynamic shape metadata but tensor shape " diff --git a/tensorflow/compiler/tf2xla/kernels/const_op.cc b/tensorflow/compiler/tf2xla/kernels/const_op.cc index a67c25be8f14f2..9d5cafe49e720c 100644 --- a/tensorflow/compiler/tf2xla/kernels/const_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/const_op.cc @@ -226,6 +226,13 @@ int64_t CountDynamicShapeContents(const TensorShapeProto& shape) { return dynamic_count; } +bool CanAttachContentsFromTensorShapeProto(const TensorShape& tensor_shape, + const TensorShapeProto& contents) { + return (tensor_shape.dims() == 0 && contents.dim_size() == 1) || + (tensor_shape.dims() == 1 && + tensor_shape.dim_size(0) == contents.dim_size()); +} + class ConstOp : public XlaOpKernel { public: explicit ConstOp(OpKernelConstruction* ctx) : XlaOpKernel(ctx) { @@ -245,17 +252,23 @@ class ConstOp : public XlaOpKernel { bool has_dynamic = false; TensorShapeProto inferred_shape_proto; + TensorShapeProto inferred_value_contents_proto; if (GetNodeAttr(ctx->op_kernel().def(), "has_dynamic", &has_dynamic).ok() && - has_dynamic) { - if (GetNodeAttr(ctx->op_kernel().def(), "user_inferred_shape", - &inferred_shape_proto) - .ok()) { - VLOG(1) << "ConstOp recovered dynamic folded-const metadata with " - << "inferred_shape=" << inferred_shape_proto.DebugString() - << " dynamic_exprs=" - << CountDynamicShapeContents(inferred_shape_proto); - } + has_dynamic && + GetNodeAttr(ctx->op_kernel().def(), "user_inferred_shape", + &inferred_shape_proto) + .ok()) { + VLOG(1) << "ConstOp recovered dynamic folded-const metadata with " + << "inferred_shape=" << inferred_shape_proto.DebugString() + << " dynamic_exprs=" + << CountDynamicShapeContents(inferred_shape_proto); } + GetNodeAttr(ctx->op_kernel().def(), "user_inferred_value_contents", + &inferred_value_contents_proto) + .IgnoreError(); + const bool has_contents_proto = inferred_value_contents_proto.dim_size() > 0; + const TensorShapeProto& contents_proto = + has_contents_proto ? inferred_value_contents_proto : inferred_shape_proto; // To avoid blowups for large constants filled with the same value, // recognize that case and emit a scalar broadcast instead. @@ -272,14 +285,14 @@ class ConstOp : public XlaOpKernel { xla::Broadcast(value, shape.dim_sizes(), shape.get_expressions()); XlaExpression output = XlaExpression::XlaOp(broadcast, ctx->expected_output_dtype(0)); - if (has_dynamic && shape.dims() == 1 && - shape.dim_size(0) == inferred_shape_proto.dim_size()) { + if ((has_contents_proto || has_dynamic) && + CanAttachContentsFromTensorShapeProto(shape, contents_proto)) { VLOG(1) << "ConstOp attaching shape contents through broadcast fast " - << "path with " << shape.dim_size(0) + << "path with " << shape.num_elements() << " entries and dynamic_exprs=" - << CountDynamicShapeContents(inferred_shape_proto); + << CountDynamicShapeContents(contents_proto); output.set_contents( - BuildShapeContentsFromTensorShapeProto(inferred_shape_proto)); + BuildShapeContentsFromTensorShapeProto(contents_proto)); } ctx->SetOutputExpression(0, output); return; @@ -290,19 +303,19 @@ class ConstOp : public XlaOpKernel { OP_REQUIRES(ctx, tensor.FromProto(cpu_allocator(), proto_), errors::InvalidArgument("Cannot parse tensor from proto: ", proto_.DebugString())); - if (has_dynamic) { + if (has_contents_proto || has_dynamic) { VLOG(1) << "ConstOp tensor path tensor_shape=" << tensor.shape().DebugString() << " inferred_rank=" - << inferred_shape_proto.dim_size(); + << contents_proto.dim_size(); } XlaExpression output = XlaExpression::Constant(tensor); - if (has_dynamic && tensor.dims() == 1 && - tensor.dim_size(0) == inferred_shape_proto.dim_size()) { + if ((has_contents_proto || has_dynamic) && + CanAttachContentsFromTensorShapeProto(tensor.shape(), contents_proto)) { VLOG(1) << "ConstOp attaching shape contents to folded const with " - << tensor.dim_size(0) << " entries and dynamic_exprs=" - << CountDynamicShapeContents(inferred_shape_proto); + << tensor.NumElements() << " entries and dynamic_exprs=" + << CountDynamicShapeContents(contents_proto); output.set_contents( - BuildShapeContentsFromTensorShapeProto(inferred_shape_proto)); + BuildShapeContentsFromTensorShapeProto(contents_proto)); } ctx->SetOutputExpression(0, output); } diff --git a/tensorflow/core/common_runtime/constant_folding.cc b/tensorflow/core/common_runtime/constant_folding.cc index e4427877c33cce..681ea1a280b77a 100644 --- a/tensorflow/core/common_runtime/constant_folding.cc +++ b/tensorflow/core/common_runtime/constant_folding.cc @@ -53,6 +53,7 @@ namespace { const char kScopedAllocatorAttrName[] = "_scoped_allocator"; const char kXlaShapeDerivedAttrName[] = "_xla_shape_derived"; +const char kUserInferredValueContentsAttrName[] = "user_inferred_value_contents"; bool IsShapeOp(const Node* n); @@ -80,6 +81,363 @@ bool GetShapeFromDirectDynamicSource(const Node* node, GetShapeFromArgNode(node, out_shape); } +bool GetConstTensor(const Node* node, Tensor* tensor) { + if (node == nullptr || !node->IsConstant()) { + return false; + } + const TensorProto* tensor_proto; + if (!GetNodeAttr(node->attrs(), "value", &tensor_proto).ok()) { + return false; + } + DataType dtype; + if (!GetNodeAttr(node->attrs(), "dtype", &dtype).ok()) { + return false; + } + *tensor = Tensor(dtype); + return tensor->FromProto(cpu_allocator(), *tensor_proto); +} + +bool GetInputConstTensor(const Node* node, int input_index, Tensor* tensor) { + const Edge* edge; + if (!node->input_edge(input_index, &edge).ok()) { + return false; + } + return GetConstTensor(edge->src(), tensor); +} + +bool GetTensorIntValues(const Tensor& tensor, std::vector* values) { + values->clear(); + if (tensor.dims() == 0) { + values->reserve(1); + switch (tensor.dtype()) { + case DT_INT32: + values->push_back(tensor.scalar()()); + return true; + case DT_INT64: + values->push_back(tensor.scalar()()); + return true; + default: + return false; + } + } + if (tensor.dims() != 1) { + return false; + } + values->reserve(tensor.NumElements()); + switch (tensor.dtype()) { + case DT_INT32: { + auto flat = tensor.flat(); + for (int i = 0; i < flat.size(); ++i) values->push_back(flat(i)); + return true; + } + case DT_INT64: { + auto flat = tensor.flat(); + for (int i = 0; i < flat.size(); ++i) values->push_back(flat(i)); + return true; + } + default: + return false; + } +} + +void CopyContentAt(const TensorShapeProto& input_contents, int64_t index, + TensorShapeProto* output_contents) { + output_contents->add_dim()->CopyFrom(input_contents.dim(index)); + if (index < input_contents.expressions_size()) { + output_contents->add_expressions()->CopyFrom(input_contents.expressions(index)); + } +} + +void AppendScalarConstantContent(int64_t value, TensorShapeProto* output_contents) { + output_contents->add_dim()->set_size(value); +} + +void AppendScalarContentFromTensor(const Tensor& tensor, + TensorShapeProto* output_contents) { + if (tensor.dtype() == DT_INT32) { + AppendScalarConstantContent(tensor.scalar()(), output_contents); + } else { + AppendScalarConstantContent(tensor.scalar()(), output_contents); + } +} + +ExpressionProto MakeConstantExpressionProto(int64_t value) { + ExpressionProto expr; + expr.set_constant_value(value); + return expr; +} + +ExpressionProto GetContentExpressionProto(const TensorShapeProto& contents, + int64_t index) { + if (index < contents.expressions_size()) { + return contents.expressions(index); + } + return MakeConstantExpressionProto(contents.dim(index).size()); +} + +ExpressionProto MakeMulExpressionProto(ExpressionProto lhs, ExpressionProto rhs) { + ExpressionProto expr; + auto* mul = expr.mutable_mul_node(); + *mul->mutable_lhs() = std::move(lhs); + *mul->mutable_rhs() = std::move(rhs); + return expr; +} + +bool TryGetFoldedValueContents(const Node* node, int output_index, + TensorShapeProto* out_contents) { + out_contents->Clear(); + if (output_index != 0) { + return false; + } + + TensorShapeProto existing_contents; + if (GetNodeAttr(node->attrs(), kUserInferredValueContentsAttrName, + &existing_contents) + .ok()) { + *out_contents = existing_contents; + return true; + } + + bool has_dynamic = false; + TensorShapeProto user_inferred_shape; + if (GetNodeAttr(node->attrs(), "has_dynamic", &has_dynamic).ok() && + has_dynamic && + GetNodeAttr(node->attrs(), "user_inferred_shape", &user_inferred_shape) + .ok()) { + *out_contents = user_inferred_shape; + return true; + } + + if (GetShapeFromDirectDynamicSource(node, out_contents)) { + return true; + } + + auto recurse_input = [&](int input_index, + TensorShapeProto* input_contents) -> bool { + const Edge* input_edge; + if (!node->input_edge(input_index, &input_edge).ok()) { + return false; + } + return TryGetFoldedValueContents(input_edge->src(), input_edge->src_output(), + input_contents); + }; + + if (node->IsIdentity() || node->type_string() == "Cast") { + return recurse_input(0, out_contents); + } + + TensorShapeProto input_contents; + if (node->type_string() == "Reshape" && recurse_input(0, &input_contents)) { + Tensor shape_tensor; + std::vector shape_dims; + if (!GetInputConstTensor(node, 1, &shape_tensor) || + !GetTensorIntValues(shape_tensor, &shape_dims)) { + return false; + } + if (input_contents.dim_size() == 1 && shape_dims.empty()) { + CopyContentAt(input_contents, 0, out_contents); + return true; + } + if (shape_dims.size() == 1 && input_contents.dim_size() == shape_dims[0]) { + out_contents->CopyFrom(input_contents); + return true; + } + return false; + } + + if (node->type_string() == "Pack") { + for (int i = 0; i < node->num_inputs(); ++i) { + TensorShapeProto scalar_contents; + Tensor scalar_tensor; + if (recurse_input(i, &scalar_contents)) { + if (scalar_contents.dim_size() != 1) { + return false; + } + CopyContentAt(scalar_contents, 0, out_contents); + } else if (GetInputConstTensor(node, i, &scalar_tensor) && + scalar_tensor.dims() == 0 && + (scalar_tensor.dtype() == DT_INT32 || + scalar_tensor.dtype() == DT_INT64)) { + AppendScalarContentFromTensor(scalar_tensor, out_contents); + } else { + return false; + } + } + return true; + } + + if (node->type_string() == "ConcatV2") { + Tensor axis_tensor; + std::vector axis_values; + if (!GetInputConstTensor(node, node->num_inputs() - 1, &axis_tensor) || + !GetTensorIntValues(axis_tensor, &axis_values) || axis_values.size() != 1) { + return false; + } + int64_t axis = axis_values[0]; + if (axis != 0 && axis != -1) { + return false; + } + for (int i = 0; i < node->num_inputs() - 1; ++i) { + TensorShapeProto part_contents; + if (!recurse_input(i, &part_contents)) { + return false; + } + for (int64_t j = 0; j < part_contents.dim_size(); ++j) { + CopyContentAt(part_contents, j, out_contents); + } + } + return true; + } + + if ((node->type_string() == "Gather" || node->type_string() == "GatherV2") && + recurse_input(0, &input_contents)) { + Tensor indices_tensor; + std::vector indices; + if (!GetInputConstTensor(node, 1, &indices_tensor) || + !GetTensorIntValues(indices_tensor, &indices)) { + return false; + } + + int64_t axis = 0; + if (node->type_string() == "GatherV2") { + Tensor axis_tensor; + std::vector axis_values; + if (!GetInputConstTensor(node, 2, &axis_tensor) || + !GetTensorIntValues(axis_tensor, &axis_values) || + axis_values.size() != 1) { + return false; + } + axis = axis_values[0]; + } + + const int64_t params_rank = 1; + if (axis < 0) axis += params_rank; + if (axis != 0) { + return false; + } + + const int64_t rank = input_contents.dim_size(); + for (int64_t index : indices) { + if (index < 0) index += rank; + if (index < 0 || index >= rank) { + return false; + } + CopyContentAt(input_contents, index, out_contents); + } + return true; + } + + if (node->type_string() == "Prod" && recurse_input(0, &input_contents)) { + Tensor reduction_indices_tensor; + std::vector axes; + bool keep_dims = false; + if (!GetInputConstTensor(node, 1, &reduction_indices_tensor) || + !GetTensorIntValues(reduction_indices_tensor, &axes) || + !GetNodeAttr(node->attrs(), "keep_dims", &keep_dims).ok() || + keep_dims || axes.size() != 1 || + (axes[0] != 0 && axes[0] != -1) || input_contents.dim_size() == 0) { + return false; + } + int64_t value = 1; + ExpressionProto expr = GetContentExpressionProto(input_contents, 0); + for (int64_t i = 0; i < input_contents.dim_size(); ++i) { + value *= input_contents.dim(i).size(); + if (i > 0) { + expr = MakeMulExpressionProto(std::move(expr), + GetContentExpressionProto(input_contents, i)); + } + } + out_contents->add_dim()->set_size(value); + out_contents->add_expressions()->Swap(&expr); + return true; + } + + if (node->type_string() == "Slice" && recurse_input(0, &input_contents)) { + Tensor begin_tensor; + Tensor size_tensor; + std::vector begin; + std::vector size; + if (!GetInputConstTensor(node, 1, &begin_tensor) || + !GetInputConstTensor(node, 2, &size_tensor) || + !GetTensorIntValues(begin_tensor, &begin) || + !GetTensorIntValues(size_tensor, &size) || begin.size() != 1 || + size.size() != 1) { + return false; + } + int64_t start = begin[0]; + if (start < 0 || start > input_contents.dim_size()) { + return false; + } + int64_t length = size[0] < 0 ? input_contents.dim_size() - start : size[0]; + if (length < 0 || start + length > input_contents.dim_size()) { + return false; + } + for (int64_t i = 0; i < length; ++i) { + CopyContentAt(input_contents, start + i, out_contents); + } + return true; + } + + if (node->type_string() == "StridedSlice" && + recurse_input(0, &input_contents)) { + Tensor begin_tensor; + Tensor end_tensor; + Tensor strides_tensor; + std::vector begin; + std::vector end; + std::vector strides; + int64_t begin_mask = 0; + int64_t end_mask = 0; + int64_t ellipsis_mask = 0; + int64_t new_axis_mask = 0; + int64_t shrink_axis_mask = 0; + if (!GetInputConstTensor(node, 1, &begin_tensor) || + !GetInputConstTensor(node, 2, &end_tensor) || + !GetInputConstTensor(node, 3, &strides_tensor) || + !GetTensorIntValues(begin_tensor, &begin) || + !GetTensorIntValues(end_tensor, &end) || + !GetTensorIntValues(strides_tensor, &strides) || begin.size() != 1 || + end.size() != 1 || strides.size() != 1 || + !GetNodeAttr(node->attrs(), "begin_mask", &begin_mask).ok() || + !GetNodeAttr(node->attrs(), "end_mask", &end_mask).ok() || + !GetNodeAttr(node->attrs(), "ellipsis_mask", &ellipsis_mask).ok() || + !GetNodeAttr(node->attrs(), "new_axis_mask", &new_axis_mask).ok() || + !GetNodeAttr(node->attrs(), "shrink_axis_mask", &shrink_axis_mask).ok()) { + return false; + } + if (ellipsis_mask != 0 || new_axis_mask != 0) { + return false; + } + const int64_t rank = input_contents.dim_size(); + int64_t stride = strides[0]; + if (stride == 0) { + return false; + } + int64_t start = (begin_mask & 1) ? (stride > 0 ? 0 : rank - 1) : begin[0]; + int64_t stop = (end_mask & 1) ? (stride > 0 ? rank : -1) : end[0]; + if (start < 0) start += rank; + if (stop < 0 && !(end_mask & 1 && stride < 0)) stop += rank; + if (shrink_axis_mask & 1) { + if (start < 0 || start >= rank) { + return false; + } + CopyContentAt(input_contents, start, out_contents); + return true; + } + if (stride < 0) { + return false; + } + start = std::max(0, start); + stop = std::min(rank, stop); + for (int64_t i = start; i < stop; i += stride) { + CopyContentAt(input_contents, i, out_contents); + } + return true; + } + + return false; +} + // For stateless RNGs ops, they are pure but device-dependent. Those ops are not // constant-foldable. static absl::flat_hash_set* kBlockList = @@ -273,16 +631,13 @@ bool IsConstantFoldable( const std::function& consider, int64_t max_constant_size_in_bytes, std::unordered_map>* shape_replacement_map) { - TensorShapeProto dynamic_shape; - if (GetShapeFromDirectDynamicSource(n, &dynamic_shape)) { - VLOG(1) << "Skipping constant folding for dynamic shape-derived node " - << n->name() << " op=" << n->type_string() - << " inferred_shape=" << dynamic_shape.DebugString(); - return false; - } - if (n->attrs().FindByString(kXlaShapeDerivedAttrName) != nullptr) { - VLOG(1) << "Skipping constant folding for shape-derived node " - << n->name() << " op=" << n->type_string(); + TensorShapeProto exact_contents; + const bool has_exact_contents = TryGetFoldedValueContents(n, 0, &exact_contents); + const bool has_dynamic = + GetShapeFromDirectDynamicSource(n, &exact_contents); + const bool is_shape_derived = + n->attrs().FindByString(kXlaShapeDerivedAttrName) != nullptr; + if ((has_dynamic || is_shape_derived) && (!has_exact_contents || n->num_outputs() > 1)) { return false; } if (n->IsConstant()) { @@ -492,6 +847,8 @@ void AddShapeNodeToConstantGraph( TensorShapeProto user_inferred_shape; const bool has_dynamic = GetShapeFromDirectDynamicSource(n, &user_inferred_shape); + TensorShapeProto exact_contents; + const bool has_exact_contents = TryGetFoldedValueContents(n, 0, &exact_contents); std::vector& added = (*node_map)[n]; const string& node_name = n->name(); for (const Tensor& t : shape_replacement_map.at(n)) { @@ -505,6 +862,9 @@ void AddShapeNodeToConstantGraph( builder.Attr("has_dynamic", has_dynamic) .Attr("user_inferred_shape", user_inferred_shape); } + if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { + builder.Attr(kUserInferredValueContentsAttrName, exact_contents); + } NodeDef def; CHECK(builder.Finalize(&def).ok()); Node* constant_node; @@ -626,6 +986,24 @@ bool ReplaceTensorWithConstant( TensorShapeProto user_inferred_shape; const bool has_dynamic = GetShapeFromDirectDynamicSource(tensor.first, &user_inferred_shape); + const bool is_shape_derived = + tensor.first->attrs().FindByString(kXlaShapeDerivedAttrName) != nullptr; + if (tensor.second != 0 && (has_dynamic || is_shape_derived)) { + VLOG(1) << "Skipping replacement of " << tensor.first->name() << " :: " + << tensor.second + << " because symbolic content preservation is only supported for " + << "single-output replacements"; + return false; + } + TensorShapeProto exact_contents; + const bool has_exact_contents = + TryGetFoldedValueContents(tensor.first, tensor.second, &exact_contents); + if ((has_dynamic || is_shape_derived) && !has_exact_contents) { + VLOG(1) << "Skipping replacement of " << tensor.first->name() << " :: " + << tensor.second + << " because constant folding could not preserve symbolic contents"; + return false; + } Node* constant_node; auto builder = NodeDefBuilder(generate_new_name(graph, node_name), "Const") .Attr("dtype", constant.dtype()) @@ -634,6 +1012,10 @@ bool ReplaceTensorWithConstant( builder.Attr("has_dynamic", has_dynamic) .Attr("user_inferred_shape", user_inferred_shape); } + if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { + builder.Attr("has_dynamic", true) + .Attr(kUserInferredValueContentsAttrName, exact_contents); + } if (partition_device) { builder.Device(partition_device->name()); } From 24c0a4a60ff125c724d293067d92125f690645c6 Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Mon, 27 Apr 2026 11:48:39 +0100 Subject: [PATCH 06/15] Serialize symbolic const contents attrs --- tensorflow/compiler/jit/kernels/xla_ops.cc | 20 +++++------- .../compiler/tf2xla/kernels/const_op.cc | 31 +++++++------------ .../core/common_runtime/constant_folding.cc | 16 ++++++++-- 3 files changed, 33 insertions(+), 34 deletions(-) diff --git a/tensorflow/compiler/jit/kernels/xla_ops.cc b/tensorflow/compiler/jit/kernels/xla_ops.cc index 018e5cb9b196a1..3a2305d544c361 100644 --- a/tensorflow/compiler/jit/kernels/xla_ops.cc +++ b/tensorflow/compiler/jit/kernels/xla_ops.cc @@ -556,6 +556,8 @@ absl::Status CompileToLocalExecutable( auto inferred_shape_it = attr_map.find("user_inferred_shape"); auto inferred_contents_it = attr_map.find("user_inferred_value_contents"); + auto inferred_contents_serialized_it = + attr_map.find("user_inferred_value_contents_serialized"); bool has_dynamic = false; auto has_dynamic_it = attr_map.find("has_dynamic"); if (has_dynamic_it != attr_map.end()) { @@ -563,12 +565,16 @@ absl::Status CompileToLocalExecutable( } if (inferred_contents_it == attr_map.end() && + inferred_contents_serialized_it == attr_map.end() && (!has_dynamic || inferred_shape_it == attr_map.end())) { return; } TensorShapeProto inferred_shape_proto; - if (inferred_contents_it != attr_map.end()) { + if (inferred_contents_serialized_it != attr_map.end()) { + inferred_shape_proto.ParseFromString( + inferred_contents_serialized_it->second.s()); + } else if (inferred_contents_it != attr_map.end()) { inferred_shape_proto = inferred_contents_it->second.shape(); } else { inferred_shape_proto = inferred_shape_it->second.shape(); @@ -579,11 +585,6 @@ absl::Status CompileToLocalExecutable( arg.constant_value.NumElements() == inferred_shape.dims()) || (TensorShapeUtils::IsScalar(arg.constant_value.shape()) && inferred_shape.dims() == 1))) { - VLOG(1) << "XlaCompileOp const arg " << arg_index - << " node=" << node_name - << " has dynamic shape metadata but tensor shape " - << arg.constant_value.shape().DebugString() - << " does not match inferred rank " << inferred_shape.dims(); return; } @@ -599,18 +600,11 @@ absl::Status CompileToLocalExecutable( } else if (arg.constant_value.dtype() == DT_INT64) { expr.set_constant_value(arg.constant_value.flat()(i)); } else { - VLOG(1) << "XlaCompileOp const arg " << arg_index - << " node=" << node_name - << " has unsupported dtype for inferred shape contents: " - << DataTypeString(arg.constant_value.dtype()); arg.constant_value_expressions.clear(); return; } arg.constant_value_expressions.push_back(std::move(expr)); } - VLOG(1) << "XlaCompileOp recovered " << arg.constant_value_expressions.size() - << " constant_value_expressions for const arg " << arg_index - << " node=" << node_name << " from user_inferred_shape"; }; auto record_dynamic_dim_value = [&](int64_t dim_size, xla::DExpr expr) { if (!saw_dynamic_dim_value) { diff --git a/tensorflow/compiler/tf2xla/kernels/const_op.cc b/tensorflow/compiler/tf2xla/kernels/const_op.cc index 9d5cafe49e720c..928d26537fdba0 100644 --- a/tensorflow/compiler/tf2xla/kernels/const_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/const_op.cc @@ -253,19 +253,24 @@ class ConstOp : public XlaOpKernel { bool has_dynamic = false; TensorShapeProto inferred_shape_proto; TensorShapeProto inferred_value_contents_proto; + string inferred_value_contents_serialized; if (GetNodeAttr(ctx->op_kernel().def(), "has_dynamic", &has_dynamic).ok() && has_dynamic && GetNodeAttr(ctx->op_kernel().def(), "user_inferred_shape", &inferred_shape_proto) .ok()) { - VLOG(1) << "ConstOp recovered dynamic folded-const metadata with " - << "inferred_shape=" << inferred_shape_proto.DebugString() - << " dynamic_exprs=" - << CountDynamicShapeContents(inferred_shape_proto); } - GetNodeAttr(ctx->op_kernel().def(), "user_inferred_value_contents", - &inferred_value_contents_proto) - .IgnoreError(); + if (GetNodeAttr(ctx->op_kernel().def(), + "user_inferred_value_contents_serialized", + &inferred_value_contents_serialized) + .ok()) { + inferred_value_contents_proto.ParseFromString( + inferred_value_contents_serialized); + } else { + GetNodeAttr(ctx->op_kernel().def(), "user_inferred_value_contents", + &inferred_value_contents_proto) + .IgnoreError(); + } const bool has_contents_proto = inferred_value_contents_proto.dim_size() > 0; const TensorShapeProto& contents_proto = has_contents_proto ? inferred_value_contents_proto : inferred_shape_proto; @@ -287,10 +292,6 @@ class ConstOp : public XlaOpKernel { XlaExpression::XlaOp(broadcast, ctx->expected_output_dtype(0)); if ((has_contents_proto || has_dynamic) && CanAttachContentsFromTensorShapeProto(shape, contents_proto)) { - VLOG(1) << "ConstOp attaching shape contents through broadcast fast " - << "path with " << shape.num_elements() - << " entries and dynamic_exprs=" - << CountDynamicShapeContents(contents_proto); output.set_contents( BuildShapeContentsFromTensorShapeProto(contents_proto)); } @@ -303,17 +304,9 @@ class ConstOp : public XlaOpKernel { OP_REQUIRES(ctx, tensor.FromProto(cpu_allocator(), proto_), errors::InvalidArgument("Cannot parse tensor from proto: ", proto_.DebugString())); - if (has_contents_proto || has_dynamic) { - VLOG(1) << "ConstOp tensor path tensor_shape=" - << tensor.shape().DebugString() << " inferred_rank=" - << contents_proto.dim_size(); - } XlaExpression output = XlaExpression::Constant(tensor); if ((has_contents_proto || has_dynamic) && CanAttachContentsFromTensorShapeProto(tensor.shape(), contents_proto)) { - VLOG(1) << "ConstOp attaching shape contents to folded const with " - << tensor.NumElements() << " entries and dynamic_exprs=" - << CountDynamicShapeContents(contents_proto); output.set_contents( BuildShapeContentsFromTensorShapeProto(contents_proto)); } diff --git a/tensorflow/core/common_runtime/constant_folding.cc b/tensorflow/core/common_runtime/constant_folding.cc index 681ea1a280b77a..9a5bf9e723ea41 100644 --- a/tensorflow/core/common_runtime/constant_folding.cc +++ b/tensorflow/core/common_runtime/constant_folding.cc @@ -54,6 +54,8 @@ namespace { const char kScopedAllocatorAttrName[] = "_scoped_allocator"; const char kXlaShapeDerivedAttrName[] = "_xla_shape_derived"; const char kUserInferredValueContentsAttrName[] = "user_inferred_value_contents"; +const char kUserInferredValueContentsSerializedAttrName[] = + "user_inferred_value_contents_serialized"; bool IsShapeOp(const Node* n); @@ -190,6 +192,14 @@ bool TryGetFoldedValueContents(const Node* node, int output_index, return false; } + string serialized_contents; + if (GetNodeAttr(node->attrs(), kUserInferredValueContentsSerializedAttrName, + &serialized_contents) + .ok() && + out_contents->ParseFromString(serialized_contents)) { + return true; + } + TensorShapeProto existing_contents; if (GetNodeAttr(node->attrs(), kUserInferredValueContentsAttrName, &existing_contents) @@ -863,7 +873,8 @@ void AddShapeNodeToConstantGraph( .Attr("user_inferred_shape", user_inferred_shape); } if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { - builder.Attr(kUserInferredValueContentsAttrName, exact_contents); + builder.Attr(kUserInferredValueContentsSerializedAttrName, + exact_contents.SerializeAsString()); } NodeDef def; CHECK(builder.Finalize(&def).ok()); @@ -1014,7 +1025,8 @@ bool ReplaceTensorWithConstant( } if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { builder.Attr("has_dynamic", true) - .Attr(kUserInferredValueContentsAttrName, exact_contents); + .Attr(kUserInferredValueContentsSerializedAttrName, + exact_contents.SerializeAsString()); } if (partition_device) { builder.Device(partition_device->name()); From 4f4f042651eef68bd37ee2813f5fb5f2908bb3e4 Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Fri, 24 Apr 2026 19:06:20 +0100 Subject: [PATCH 07/15] Avoid partial non-CPU multi-output constant replacement --- tensorflow/core/common_runtime/constant_folding.cc | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/tensorflow/core/common_runtime/constant_folding.cc b/tensorflow/core/common_runtime/constant_folding.cc index 9a5bf9e723ea41..9cee31219b3c81 100644 --- a/tensorflow/core/common_runtime/constant_folding.cc +++ b/tensorflow/core/common_runtime/constant_folding.cc @@ -966,6 +966,12 @@ bool ReplaceTensorWithConstant( ? DeviceType{partition_device->device_type()} : DEVICE_CPU; if (partition_device && device_type != DEVICE_CPU) { + // Constant folding replaces one output edge-set at a time. Be + // conservative for non-CPU multi-output ops, since partially replacing a + // node can violate per-output placement or memory-type assumptions. + if (tensor.first->num_outputs() > 1) { + return false; + } MemoryTypeVector input_mvec; MemoryTypeVector output_mvec; if (!MemoryTypesForNode(graph->op_registry(), device_type, From 6035005fa19df56c65f6a1f75ecf549a4cf5ced1 Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Mon, 27 Apr 2026 11:38:23 +0100 Subject: [PATCH 08/15] Add symbolic constant folding coverage --- tensorflow/core/common_runtime/BUILD | 4 + .../common_runtime/constant_folding_test.cc | 201 ++++++++++++++++++ 2 files changed, 205 insertions(+) diff --git a/tensorflow/core/common_runtime/BUILD b/tensorflow/core/common_runtime/BUILD index 301015eba61fe8..24089b11472b6b 100644 --- a/tensorflow/core/common_runtime/BUILD +++ b/tensorflow/core/common_runtime/BUILD @@ -2754,6 +2754,7 @@ tf_cc_test( ":direct_session_internal", "//tensorflow/cc:cc_ops", "//tensorflow/cc:cc_ops_internal", + "//tensorflow/cc:function_ops", "//tensorflow/cc:sendrecv_ops", "//tensorflow/core:framework", "//tensorflow/core:framework_internal", @@ -2769,9 +2770,12 @@ tf_cc_test( "//tensorflow/core/kernels:cast_op", "//tensorflow/core/kernels:concat_op", "//tensorflow/core/kernels:cwise_op", + "//tensorflow/core/kernels:gather_op", "//tensorflow/core/kernels:identity_op", "//tensorflow/core/kernels:immutable_constant_op", "//tensorflow/core/kernels:matmul_op", + "//tensorflow/core/kernels:reshape_op", + "//tensorflow/core/kernels:slice_op", "//tensorflow/core/kernels:topk_op", "@eigen_archive//:eigen3", ], diff --git a/tensorflow/core/common_runtime/constant_folding_test.cc b/tensorflow/core/common_runtime/constant_folding_test.cc index 481a85add4893c..84b5f324f9a62c 100644 --- a/tensorflow/core/common_runtime/constant_folding_test.cc +++ b/tensorflow/core/common_runtime/constant_folding_test.cc @@ -21,6 +21,7 @@ limitations under the License. #include #include "tensorflow/cc/ops/array_ops_internal.h" +#include "tensorflow/cc/ops/function_ops.h" #include "tensorflow/cc/ops/nn_ops.h" #include "tensorflow/cc/ops/sendrecv_ops.h" #include "tensorflow/cc/ops/standard_ops.h" @@ -32,6 +33,7 @@ limitations under the License. #include "tensorflow/core/framework/node_def_util.h" #include "tensorflow/core/framework/tensor.h" #include "tensorflow/core/framework/tensor_shape.h" +#include "tensorflow/core/framework/tensor_shape_expr.h" #include "tensorflow/core/framework/tensor_testutil.h" #include "tensorflow/core/framework/types.h" #include "tensorflow/core/graph/node_builder.h" @@ -45,6 +47,15 @@ limitations under the License. namespace tensorflow { namespace { +TensorShapeProto MakeDynamicShapeProto977x16() { + TensorShapeProto proto; + proto.add_dim()->set_size(977); + proto.add_dim()->set_size(16); + proto.add_expressions()->set_variable_id(0); + proto.add_expressions()->set_constant_value(16); + return proto; +} + class ConstantFoldingTest : public ::testing::Test { protected: template @@ -634,6 +645,196 @@ TEST_F(ConstantFoldingTest, ConstShapeKnown) { } } +TEST_F(ConstantFoldingTest, FoldShapeFromDynamicArgPreservesContents) { + Graph g(OpRegistry::Global()); + Scope s = Scope::NewRootScope(); + auto arg = ops::_Arg(s.WithOpName("arg"), DT_FLOAT, 0); + auto shape = ops::Shape(s.WithOpName("shape"), arg); + auto send = ops::_Send(s.WithOpName("send"), shape, "send", "sender", 0, + "receiver"); + TF_ASSERT_OK(s.ToGraph(&g)); + + std::unordered_map index_by_name = g.BuildNodeNameIndex(); + Node* arg_node = index_by_name.at("arg"); + arg_node->AddAttr("_output_shapes", + std::vector{MakeDynamicShapeProto977x16()}); + + PartialTensorShape partial_shape({977, 16}); + std::unordered_map> shape_map; + shape_map[arg_node->name()].push_back(partial_shape); + + ConstantFoldingOptions opts; + opts.shape_map = &shape_map; + bool was_mutated = false; + TF_ASSERT_OK( + ConstantFold(opts, nullptr, Env::Default(), nullptr, &g, &was_mutated)); + + index_by_name = g.BuildNodeNameIndex(); + Node* send_node = index_by_name.at("send"); + const Edge* send_input = nullptr; + TF_ASSERT_OK(send_node->input_edge(0, &send_input)); + Node* folded = send_input->src(); + ExpectNodeEqual(folded, {977, 16}, {2}); + + // This test checks the stable contract we care about: after folding + // Shape(arg), the replacement Const still carries symbolic contents. + string serialized_contents_proto; + TF_ASSERT_OK(GetNodeAttr(folded->attrs(), + "user_inferred_value_contents_serialized", + &serialized_contents_proto)); + TensorShapeProto contents_proto; + ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); + ASSERT_EQ(contents_proto.expressions_size(), 2); + EXPECT_TRUE(IsDynamicDimExpr(contents_proto.expressions(0))); + EXPECT_FALSE(IsDynamicDimExpr(contents_proto.expressions(1))); + EXPECT_EQ(contents_proto.dim(0).size(), 977); + EXPECT_EQ(contents_proto.dim(1).size(), 16); +} + +TEST_F(ConstantFoldingTest, FoldSliceOfDynamicShapePreservesContents) { + Graph g(OpRegistry::Global()); + Scope s = Scope::NewRootScope(); + auto arg = ops::_Arg(s.WithOpName("arg"), DT_FLOAT, 0); + auto shape = ops::Shape(s.WithOpName("shape"), arg); + auto begin = ops::Const(s.WithOpName("begin"), {0}, {1}); + auto size = ops::Const(s.WithOpName("size"), {1}, {1}); + auto slice = ops::Slice(s.WithOpName("slice"), shape, begin, size); + auto send = ops::_Send(s.WithOpName("send"), slice, "send", "sender", 0, + "receiver"); + TF_ASSERT_OK(s.ToGraph(&g)); + + std::unordered_map index_by_name = g.BuildNodeNameIndex(); + Node* arg_node = index_by_name.at("arg"); + arg_node->AddAttr("_output_shapes", + std::vector{MakeDynamicShapeProto977x16()}); + + PartialTensorShape partial_shape({977, 16}); + std::unordered_map> shape_map; + shape_map[arg_node->name()].push_back(partial_shape); + + ConstantFoldingOptions opts; + opts.shape_map = &shape_map; + bool was_mutated = false; + TF_ASSERT_OK( + ConstantFold(opts, nullptr, Env::Default(), nullptr, &g, &was_mutated)); + + index_by_name = g.BuildNodeNameIndex(); + Node* send_node = index_by_name.at("send"); + const Edge* send_input = nullptr; + TF_ASSERT_OK(send_node->input_edge(0, &send_input)); + Node* folded = send_input->src(); + ExpectNodeEqual(folded, {977}, {1}); + + // Folding Slice(Shape(arg), [0], [1]) should preserve the selected symbolic + // content, not just the concrete value 977. + string serialized_contents_proto; + TF_ASSERT_OK(GetNodeAttr(folded->attrs(), + "user_inferred_value_contents_serialized", + &serialized_contents_proto)); + TensorShapeProto contents_proto; + ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); + ASSERT_EQ(contents_proto.expressions_size(), 1); + EXPECT_TRUE(IsDynamicDimExpr(contents_proto.expressions(0))); + EXPECT_EQ(contents_proto.dim(0).size(), 977); +} + +TEST_F(ConstantFoldingTest, FoldGatherOfDynamicShapePreservesContents) { + Graph g(OpRegistry::Global()); + Scope s = Scope::NewRootScope(); + auto arg = ops::_Arg(s.WithOpName("arg"), DT_FLOAT, 0); + auto shape = ops::Shape(s.WithOpName("shape"), arg); + auto index = ops::Const(s.WithOpName("index"), 0); + auto gather = ops::GatherV2(s.WithOpName("gather"), shape, index, + ops::Const(s.WithOpName("axis"), 0)); + auto send = ops::_Send(s.WithOpName("send"), gather, "send", "sender", 0, + "receiver"); + TF_ASSERT_OK(s.ToGraph(&g)); + + std::unordered_map index_by_name = g.BuildNodeNameIndex(); + Node* arg_node = index_by_name.at("arg"); + arg_node->AddAttr("_output_shapes", + std::vector{MakeDynamicShapeProto977x16()}); + + PartialTensorShape partial_shape({977, 16}); + std::unordered_map> shape_map; + shape_map[arg_node->name()].push_back(partial_shape); + + ConstantFoldingOptions opts; + opts.shape_map = &shape_map; + bool was_mutated = false; + TF_ASSERT_OK( + ConstantFold(opts, nullptr, Env::Default(), nullptr, &g, &was_mutated)); + + index_by_name = g.BuildNodeNameIndex(); + Node* send_node = index_by_name.at("send"); + const Edge* send_input = nullptr; + TF_ASSERT_OK(send_node->input_edge(0, &send_input)); + Node* folded = send_input->src(); + ExpectNodeEqual(folded, {977}, {}); + + // Folding Gather(Shape(arg), 0, axis=0) should preserve the selected + // symbolic content, not just the concrete scalar value 977. + string serialized_contents_proto; + TF_ASSERT_OK(GetNodeAttr(folded->attrs(), + "user_inferred_value_contents_serialized", + &serialized_contents_proto)); + TensorShapeProto contents_proto; + ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); + ASSERT_EQ(contents_proto.expressions_size(), 1); + EXPECT_TRUE(IsDynamicDimExpr(contents_proto.expressions(0))); + EXPECT_EQ(contents_proto.dim(0).size(), 977); +} + +TEST_F(ConstantFoldingTest, FoldReshapeOfDynamicShapePreservesContents) { + Graph g(OpRegistry::Global()); + Scope s = Scope::NewRootScope(); + auto arg = ops::_Arg(s.WithOpName("arg"), DT_FLOAT, 0); + auto shape = ops::Shape(s.WithOpName("shape"), arg); + auto begin = ops::Const(s.WithOpName("begin"), {0}, {1}); + auto size = ops::Const(s.WithOpName("size"), {1}, {1}); + auto slice = ops::Slice(s.WithOpName("slice"), shape, begin, size); + auto scalar_shape = ops::Const(s.WithOpName("scalar_shape"), {}); + auto reshape = + ops::Reshape(s.WithOpName("reshape"), slice, scalar_shape); + auto send = ops::_Send(s.WithOpName("send"), reshape, "send", "sender", 0, + "receiver"); + TF_ASSERT_OK(s.ToGraph(&g)); + + std::unordered_map index_by_name = g.BuildNodeNameIndex(); + Node* arg_node = index_by_name.at("arg"); + arg_node->AddAttr("_output_shapes", + std::vector{MakeDynamicShapeProto977x16()}); + + PartialTensorShape partial_shape({977, 16}); + std::unordered_map> shape_map; + shape_map[arg_node->name()].push_back(partial_shape); + + ConstantFoldingOptions opts; + opts.shape_map = &shape_map; + bool was_mutated = false; + TF_ASSERT_OK( + ConstantFold(opts, nullptr, Env::Default(), nullptr, &g, &was_mutated)); + + index_by_name = g.BuildNodeNameIndex(); + Node* send_node = index_by_name.at("send"); + const Edge* send_input = nullptr; + TF_ASSERT_OK(send_node->input_edge(0, &send_input)); + Node* folded = send_input->src(); + ExpectNodeEqual(folded, {977}, {}); + + // Folding Reshape(Slice(Shape(arg), [0], [1]), []) should preserve the + // selected symbolic content when the shape-vector result is scalarized. + string serialized_contents_proto; + TF_ASSERT_OK(GetNodeAttr(folded->attrs(), + "user_inferred_value_contents_serialized", + &serialized_contents_proto)); + TensorShapeProto contents_proto; + ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); + ASSERT_EQ(contents_proto.expressions_size(), 1); + EXPECT_TRUE(IsDynamicDimExpr(contents_proto.expressions(0))); + EXPECT_EQ(contents_proto.dim(0).size(), 977); +} + TEST_F(ConstantFoldingTest, NoReplacePartialOutput) { Graph g(OpRegistry::Global()); { From 7e71c2bb7e7d62dff6ada07e1da42f9d4c8d3dc3 Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Mon, 27 Apr 2026 12:08:05 +0100 Subject: [PATCH 09/15] Revert "Serialize symbolic const contents attrs" This reverts commit 56b8330ef5a6a2de0109dc077aceda54efad91c3. --- tensorflow/compiler/jit/kernels/xla_ops.cc | 20 +++++++----- .../compiler/tf2xla/kernels/const_op.cc | 31 ++++++++++++------- .../core/common_runtime/constant_folding.cc | 16 ++-------- 3 files changed, 34 insertions(+), 33 deletions(-) diff --git a/tensorflow/compiler/jit/kernels/xla_ops.cc b/tensorflow/compiler/jit/kernels/xla_ops.cc index 3a2305d544c361..018e5cb9b196a1 100644 --- a/tensorflow/compiler/jit/kernels/xla_ops.cc +++ b/tensorflow/compiler/jit/kernels/xla_ops.cc @@ -556,8 +556,6 @@ absl::Status CompileToLocalExecutable( auto inferred_shape_it = attr_map.find("user_inferred_shape"); auto inferred_contents_it = attr_map.find("user_inferred_value_contents"); - auto inferred_contents_serialized_it = - attr_map.find("user_inferred_value_contents_serialized"); bool has_dynamic = false; auto has_dynamic_it = attr_map.find("has_dynamic"); if (has_dynamic_it != attr_map.end()) { @@ -565,16 +563,12 @@ absl::Status CompileToLocalExecutable( } if (inferred_contents_it == attr_map.end() && - inferred_contents_serialized_it == attr_map.end() && (!has_dynamic || inferred_shape_it == attr_map.end())) { return; } TensorShapeProto inferred_shape_proto; - if (inferred_contents_serialized_it != attr_map.end()) { - inferred_shape_proto.ParseFromString( - inferred_contents_serialized_it->second.s()); - } else if (inferred_contents_it != attr_map.end()) { + if (inferred_contents_it != attr_map.end()) { inferred_shape_proto = inferred_contents_it->second.shape(); } else { inferred_shape_proto = inferred_shape_it->second.shape(); @@ -585,6 +579,11 @@ absl::Status CompileToLocalExecutable( arg.constant_value.NumElements() == inferred_shape.dims()) || (TensorShapeUtils::IsScalar(arg.constant_value.shape()) && inferred_shape.dims() == 1))) { + VLOG(1) << "XlaCompileOp const arg " << arg_index + << " node=" << node_name + << " has dynamic shape metadata but tensor shape " + << arg.constant_value.shape().DebugString() + << " does not match inferred rank " << inferred_shape.dims(); return; } @@ -600,11 +599,18 @@ absl::Status CompileToLocalExecutable( } else if (arg.constant_value.dtype() == DT_INT64) { expr.set_constant_value(arg.constant_value.flat()(i)); } else { + VLOG(1) << "XlaCompileOp const arg " << arg_index + << " node=" << node_name + << " has unsupported dtype for inferred shape contents: " + << DataTypeString(arg.constant_value.dtype()); arg.constant_value_expressions.clear(); return; } arg.constant_value_expressions.push_back(std::move(expr)); } + VLOG(1) << "XlaCompileOp recovered " << arg.constant_value_expressions.size() + << " constant_value_expressions for const arg " << arg_index + << " node=" << node_name << " from user_inferred_shape"; }; auto record_dynamic_dim_value = [&](int64_t dim_size, xla::DExpr expr) { if (!saw_dynamic_dim_value) { diff --git a/tensorflow/compiler/tf2xla/kernels/const_op.cc b/tensorflow/compiler/tf2xla/kernels/const_op.cc index 928d26537fdba0..9d5cafe49e720c 100644 --- a/tensorflow/compiler/tf2xla/kernels/const_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/const_op.cc @@ -253,24 +253,19 @@ class ConstOp : public XlaOpKernel { bool has_dynamic = false; TensorShapeProto inferred_shape_proto; TensorShapeProto inferred_value_contents_proto; - string inferred_value_contents_serialized; if (GetNodeAttr(ctx->op_kernel().def(), "has_dynamic", &has_dynamic).ok() && has_dynamic && GetNodeAttr(ctx->op_kernel().def(), "user_inferred_shape", &inferred_shape_proto) .ok()) { + VLOG(1) << "ConstOp recovered dynamic folded-const metadata with " + << "inferred_shape=" << inferred_shape_proto.DebugString() + << " dynamic_exprs=" + << CountDynamicShapeContents(inferred_shape_proto); } - if (GetNodeAttr(ctx->op_kernel().def(), - "user_inferred_value_contents_serialized", - &inferred_value_contents_serialized) - .ok()) { - inferred_value_contents_proto.ParseFromString( - inferred_value_contents_serialized); - } else { - GetNodeAttr(ctx->op_kernel().def(), "user_inferred_value_contents", - &inferred_value_contents_proto) - .IgnoreError(); - } + GetNodeAttr(ctx->op_kernel().def(), "user_inferred_value_contents", + &inferred_value_contents_proto) + .IgnoreError(); const bool has_contents_proto = inferred_value_contents_proto.dim_size() > 0; const TensorShapeProto& contents_proto = has_contents_proto ? inferred_value_contents_proto : inferred_shape_proto; @@ -292,6 +287,10 @@ class ConstOp : public XlaOpKernel { XlaExpression::XlaOp(broadcast, ctx->expected_output_dtype(0)); if ((has_contents_proto || has_dynamic) && CanAttachContentsFromTensorShapeProto(shape, contents_proto)) { + VLOG(1) << "ConstOp attaching shape contents through broadcast fast " + << "path with " << shape.num_elements() + << " entries and dynamic_exprs=" + << CountDynamicShapeContents(contents_proto); output.set_contents( BuildShapeContentsFromTensorShapeProto(contents_proto)); } @@ -304,9 +303,17 @@ class ConstOp : public XlaOpKernel { OP_REQUIRES(ctx, tensor.FromProto(cpu_allocator(), proto_), errors::InvalidArgument("Cannot parse tensor from proto: ", proto_.DebugString())); + if (has_contents_proto || has_dynamic) { + VLOG(1) << "ConstOp tensor path tensor_shape=" + << tensor.shape().DebugString() << " inferred_rank=" + << contents_proto.dim_size(); + } XlaExpression output = XlaExpression::Constant(tensor); if ((has_contents_proto || has_dynamic) && CanAttachContentsFromTensorShapeProto(tensor.shape(), contents_proto)) { + VLOG(1) << "ConstOp attaching shape contents to folded const with " + << tensor.NumElements() << " entries and dynamic_exprs=" + << CountDynamicShapeContents(contents_proto); output.set_contents( BuildShapeContentsFromTensorShapeProto(contents_proto)); } diff --git a/tensorflow/core/common_runtime/constant_folding.cc b/tensorflow/core/common_runtime/constant_folding.cc index 9cee31219b3c81..c2c2eed6261af2 100644 --- a/tensorflow/core/common_runtime/constant_folding.cc +++ b/tensorflow/core/common_runtime/constant_folding.cc @@ -54,8 +54,6 @@ namespace { const char kScopedAllocatorAttrName[] = "_scoped_allocator"; const char kXlaShapeDerivedAttrName[] = "_xla_shape_derived"; const char kUserInferredValueContentsAttrName[] = "user_inferred_value_contents"; -const char kUserInferredValueContentsSerializedAttrName[] = - "user_inferred_value_contents_serialized"; bool IsShapeOp(const Node* n); @@ -192,14 +190,6 @@ bool TryGetFoldedValueContents(const Node* node, int output_index, return false; } - string serialized_contents; - if (GetNodeAttr(node->attrs(), kUserInferredValueContentsSerializedAttrName, - &serialized_contents) - .ok() && - out_contents->ParseFromString(serialized_contents)) { - return true; - } - TensorShapeProto existing_contents; if (GetNodeAttr(node->attrs(), kUserInferredValueContentsAttrName, &existing_contents) @@ -873,8 +863,7 @@ void AddShapeNodeToConstantGraph( .Attr("user_inferred_shape", user_inferred_shape); } if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { - builder.Attr(kUserInferredValueContentsSerializedAttrName, - exact_contents.SerializeAsString()); + builder.Attr(kUserInferredValueContentsAttrName, exact_contents); } NodeDef def; CHECK(builder.Finalize(&def).ok()); @@ -1031,8 +1020,7 @@ bool ReplaceTensorWithConstant( } if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { builder.Attr("has_dynamic", true) - .Attr(kUserInferredValueContentsSerializedAttrName, - exact_contents.SerializeAsString()); + .Attr(kUserInferredValueContentsAttrName, exact_contents); } if (partition_device) { builder.Device(partition_device->name()); From e501ec31ac75124a6344af38c78b7276cb2f45d8 Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Mon, 27 Apr 2026 12:26:41 +0100 Subject: [PATCH 10/15] Reapply "Serialize symbolic const contents attrs" This reverts commit ad7ecc23db86e1873def6bcb4100177e91584ebf. --- tensorflow/compiler/jit/kernels/xla_ops.cc | 20 +++++------- .../compiler/tf2xla/kernels/const_op.cc | 31 +++++++------------ .../core/common_runtime/constant_folding.cc | 16 ++++++++-- 3 files changed, 33 insertions(+), 34 deletions(-) diff --git a/tensorflow/compiler/jit/kernels/xla_ops.cc b/tensorflow/compiler/jit/kernels/xla_ops.cc index 018e5cb9b196a1..3a2305d544c361 100644 --- a/tensorflow/compiler/jit/kernels/xla_ops.cc +++ b/tensorflow/compiler/jit/kernels/xla_ops.cc @@ -556,6 +556,8 @@ absl::Status CompileToLocalExecutable( auto inferred_shape_it = attr_map.find("user_inferred_shape"); auto inferred_contents_it = attr_map.find("user_inferred_value_contents"); + auto inferred_contents_serialized_it = + attr_map.find("user_inferred_value_contents_serialized"); bool has_dynamic = false; auto has_dynamic_it = attr_map.find("has_dynamic"); if (has_dynamic_it != attr_map.end()) { @@ -563,12 +565,16 @@ absl::Status CompileToLocalExecutable( } if (inferred_contents_it == attr_map.end() && + inferred_contents_serialized_it == attr_map.end() && (!has_dynamic || inferred_shape_it == attr_map.end())) { return; } TensorShapeProto inferred_shape_proto; - if (inferred_contents_it != attr_map.end()) { + if (inferred_contents_serialized_it != attr_map.end()) { + inferred_shape_proto.ParseFromString( + inferred_contents_serialized_it->second.s()); + } else if (inferred_contents_it != attr_map.end()) { inferred_shape_proto = inferred_contents_it->second.shape(); } else { inferred_shape_proto = inferred_shape_it->second.shape(); @@ -579,11 +585,6 @@ absl::Status CompileToLocalExecutable( arg.constant_value.NumElements() == inferred_shape.dims()) || (TensorShapeUtils::IsScalar(arg.constant_value.shape()) && inferred_shape.dims() == 1))) { - VLOG(1) << "XlaCompileOp const arg " << arg_index - << " node=" << node_name - << " has dynamic shape metadata but tensor shape " - << arg.constant_value.shape().DebugString() - << " does not match inferred rank " << inferred_shape.dims(); return; } @@ -599,18 +600,11 @@ absl::Status CompileToLocalExecutable( } else if (arg.constant_value.dtype() == DT_INT64) { expr.set_constant_value(arg.constant_value.flat()(i)); } else { - VLOG(1) << "XlaCompileOp const arg " << arg_index - << " node=" << node_name - << " has unsupported dtype for inferred shape contents: " - << DataTypeString(arg.constant_value.dtype()); arg.constant_value_expressions.clear(); return; } arg.constant_value_expressions.push_back(std::move(expr)); } - VLOG(1) << "XlaCompileOp recovered " << arg.constant_value_expressions.size() - << " constant_value_expressions for const arg " << arg_index - << " node=" << node_name << " from user_inferred_shape"; }; auto record_dynamic_dim_value = [&](int64_t dim_size, xla::DExpr expr) { if (!saw_dynamic_dim_value) { diff --git a/tensorflow/compiler/tf2xla/kernels/const_op.cc b/tensorflow/compiler/tf2xla/kernels/const_op.cc index 9d5cafe49e720c..928d26537fdba0 100644 --- a/tensorflow/compiler/tf2xla/kernels/const_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/const_op.cc @@ -253,19 +253,24 @@ class ConstOp : public XlaOpKernel { bool has_dynamic = false; TensorShapeProto inferred_shape_proto; TensorShapeProto inferred_value_contents_proto; + string inferred_value_contents_serialized; if (GetNodeAttr(ctx->op_kernel().def(), "has_dynamic", &has_dynamic).ok() && has_dynamic && GetNodeAttr(ctx->op_kernel().def(), "user_inferred_shape", &inferred_shape_proto) .ok()) { - VLOG(1) << "ConstOp recovered dynamic folded-const metadata with " - << "inferred_shape=" << inferred_shape_proto.DebugString() - << " dynamic_exprs=" - << CountDynamicShapeContents(inferred_shape_proto); } - GetNodeAttr(ctx->op_kernel().def(), "user_inferred_value_contents", - &inferred_value_contents_proto) - .IgnoreError(); + if (GetNodeAttr(ctx->op_kernel().def(), + "user_inferred_value_contents_serialized", + &inferred_value_contents_serialized) + .ok()) { + inferred_value_contents_proto.ParseFromString( + inferred_value_contents_serialized); + } else { + GetNodeAttr(ctx->op_kernel().def(), "user_inferred_value_contents", + &inferred_value_contents_proto) + .IgnoreError(); + } const bool has_contents_proto = inferred_value_contents_proto.dim_size() > 0; const TensorShapeProto& contents_proto = has_contents_proto ? inferred_value_contents_proto : inferred_shape_proto; @@ -287,10 +292,6 @@ class ConstOp : public XlaOpKernel { XlaExpression::XlaOp(broadcast, ctx->expected_output_dtype(0)); if ((has_contents_proto || has_dynamic) && CanAttachContentsFromTensorShapeProto(shape, contents_proto)) { - VLOG(1) << "ConstOp attaching shape contents through broadcast fast " - << "path with " << shape.num_elements() - << " entries and dynamic_exprs=" - << CountDynamicShapeContents(contents_proto); output.set_contents( BuildShapeContentsFromTensorShapeProto(contents_proto)); } @@ -303,17 +304,9 @@ class ConstOp : public XlaOpKernel { OP_REQUIRES(ctx, tensor.FromProto(cpu_allocator(), proto_), errors::InvalidArgument("Cannot parse tensor from proto: ", proto_.DebugString())); - if (has_contents_proto || has_dynamic) { - VLOG(1) << "ConstOp tensor path tensor_shape=" - << tensor.shape().DebugString() << " inferred_rank=" - << contents_proto.dim_size(); - } XlaExpression output = XlaExpression::Constant(tensor); if ((has_contents_proto || has_dynamic) && CanAttachContentsFromTensorShapeProto(tensor.shape(), contents_proto)) { - VLOG(1) << "ConstOp attaching shape contents to folded const with " - << tensor.NumElements() << " entries and dynamic_exprs=" - << CountDynamicShapeContents(contents_proto); output.set_contents( BuildShapeContentsFromTensorShapeProto(contents_proto)); } diff --git a/tensorflow/core/common_runtime/constant_folding.cc b/tensorflow/core/common_runtime/constant_folding.cc index c2c2eed6261af2..9cee31219b3c81 100644 --- a/tensorflow/core/common_runtime/constant_folding.cc +++ b/tensorflow/core/common_runtime/constant_folding.cc @@ -54,6 +54,8 @@ namespace { const char kScopedAllocatorAttrName[] = "_scoped_allocator"; const char kXlaShapeDerivedAttrName[] = "_xla_shape_derived"; const char kUserInferredValueContentsAttrName[] = "user_inferred_value_contents"; +const char kUserInferredValueContentsSerializedAttrName[] = + "user_inferred_value_contents_serialized"; bool IsShapeOp(const Node* n); @@ -190,6 +192,14 @@ bool TryGetFoldedValueContents(const Node* node, int output_index, return false; } + string serialized_contents; + if (GetNodeAttr(node->attrs(), kUserInferredValueContentsSerializedAttrName, + &serialized_contents) + .ok() && + out_contents->ParseFromString(serialized_contents)) { + return true; + } + TensorShapeProto existing_contents; if (GetNodeAttr(node->attrs(), kUserInferredValueContentsAttrName, &existing_contents) @@ -863,7 +873,8 @@ void AddShapeNodeToConstantGraph( .Attr("user_inferred_shape", user_inferred_shape); } if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { - builder.Attr(kUserInferredValueContentsAttrName, exact_contents); + builder.Attr(kUserInferredValueContentsSerializedAttrName, + exact_contents.SerializeAsString()); } NodeDef def; CHECK(builder.Finalize(&def).ok()); @@ -1020,7 +1031,8 @@ bool ReplaceTensorWithConstant( } if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { builder.Attr("has_dynamic", true) - .Attr(kUserInferredValueContentsAttrName, exact_contents); + .Attr(kUserInferredValueContentsSerializedAttrName, + exact_contents.SerializeAsString()); } if (partition_device) { builder.Device(partition_device->name()); From 98057078d550403f22b1dbde69a0a1463fab19d4 Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Mon, 27 Apr 2026 13:28:55 +0100 Subject: [PATCH 11/15] Address symbolic constant folding review feedback --- tensorflow/compiler/jit/kernels/xla_ops.cc | 26 +++++- .../compiler/tf2xla/kernels/const_op.cc | 29 +++++- .../core/common_runtime/constant_folding.cc | 93 ++++++++++++++++--- .../common_runtime/constant_folding_test.cc | 8 +- 4 files changed, 132 insertions(+), 24 deletions(-) diff --git a/tensorflow/compiler/jit/kernels/xla_ops.cc b/tensorflow/compiler/jit/kernels/xla_ops.cc index 3a2305d544c361..f60f42c3474de8 100644 --- a/tensorflow/compiler/jit/kernels/xla_ops.cc +++ b/tensorflow/compiler/jit/kernels/xla_ops.cc @@ -101,6 +101,14 @@ limitations under the License. namespace tensorflow { namespace { +constexpr char kUserInferredValueContentsAttrName[] = + "_user_inferred_value_contents"; +constexpr char kUserInferredValueContentsSerializedAttrName[] = + "_user_inferred_value_contents_serialized"; +constexpr char kLegacyUserInferredValueContentsAttrName[] = + "user_inferred_value_contents"; +constexpr char kLegacyUserInferredValueContentsSerializedAttrName[] = + "user_inferred_value_contents_serialized"; using XlaDeviceCompiler = DeviceCompiler; using PjRtDeviceCompiler = @@ -555,9 +563,17 @@ absl::Status CompileToLocalExecutable( auto inferred_shape_it = attr_map.find("user_inferred_shape"); auto inferred_contents_it = - attr_map.find("user_inferred_value_contents"); + attr_map.find(kUserInferredValueContentsAttrName); + if (inferred_contents_it == attr_map.end()) { + inferred_contents_it = + attr_map.find(kLegacyUserInferredValueContentsAttrName); + } auto inferred_contents_serialized_it = - attr_map.find("user_inferred_value_contents_serialized"); + attr_map.find(kUserInferredValueContentsSerializedAttrName); + if (inferred_contents_serialized_it == attr_map.end()) { + inferred_contents_serialized_it = + attr_map.find(kLegacyUserInferredValueContentsSerializedAttrName); + } bool has_dynamic = false; auto has_dynamic_it = attr_map.find("has_dynamic"); if (has_dynamic_it != attr_map.end()) { @@ -572,8 +588,10 @@ absl::Status CompileToLocalExecutable( TensorShapeProto inferred_shape_proto; if (inferred_contents_serialized_it != attr_map.end()) { - inferred_shape_proto.ParseFromString( - inferred_contents_serialized_it->second.s()); + if (!inferred_shape_proto.ParseFromString( + inferred_contents_serialized_it->second.s())) { + return; + } } else if (inferred_contents_it != attr_map.end()) { inferred_shape_proto = inferred_contents_it->second.shape(); } else { diff --git a/tensorflow/compiler/tf2xla/kernels/const_op.cc b/tensorflow/compiler/tf2xla/kernels/const_op.cc index 928d26537fdba0..254f2c26dcb374 100644 --- a/tensorflow/compiler/tf2xla/kernels/const_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/const_op.cc @@ -35,6 +35,15 @@ limitations under the License. namespace tensorflow { namespace { +constexpr char kUserInferredValueContentsAttrName[] = + "_user_inferred_value_contents"; +constexpr char kUserInferredValueContentsSerializedAttrName[] = + "_user_inferred_value_contents_serialized"; +constexpr char kLegacyUserInferredValueContentsAttrName[] = + "user_inferred_value_contents"; +constexpr char kLegacyUserInferredValueContentsSerializedAttrName[] = + "user_inferred_value_contents_serialized"; + template DstT CastTo(SrcT src) { return static_cast(src); @@ -261,15 +270,27 @@ class ConstOp : public XlaOpKernel { .ok()) { } if (GetNodeAttr(ctx->op_kernel().def(), - "user_inferred_value_contents_serialized", + kUserInferredValueContentsSerializedAttrName, + &inferred_value_contents_serialized) + .ok() || + GetNodeAttr(ctx->op_kernel().def(), + kLegacyUserInferredValueContentsSerializedAttrName, &inferred_value_contents_serialized) .ok()) { - inferred_value_contents_proto.ParseFromString( - inferred_value_contents_serialized); + if (!inferred_value_contents_proto.ParseFromString( + inferred_value_contents_serialized)) { + inferred_value_contents_proto.Clear(); + } } else { - GetNodeAttr(ctx->op_kernel().def(), "user_inferred_value_contents", + GetNodeAttr(ctx->op_kernel().def(), kUserInferredValueContentsAttrName, &inferred_value_contents_proto) .IgnoreError(); + if (inferred_value_contents_proto.dim_size() == 0) { + GetNodeAttr(ctx->op_kernel().def(), + kLegacyUserInferredValueContentsAttrName, + &inferred_value_contents_proto) + .IgnoreError(); + } } const bool has_contents_proto = inferred_value_contents_proto.dim_size() > 0; const TensorShapeProto& contents_proto = diff --git a/tensorflow/core/common_runtime/constant_folding.cc b/tensorflow/core/common_runtime/constant_folding.cc index 9cee31219b3c81..4bbcd04db68863 100644 --- a/tensorflow/core/common_runtime/constant_folding.cc +++ b/tensorflow/core/common_runtime/constant_folding.cc @@ -53,8 +53,13 @@ namespace { const char kScopedAllocatorAttrName[] = "_scoped_allocator"; const char kXlaShapeDerivedAttrName[] = "_xla_shape_derived"; -const char kUserInferredValueContentsAttrName[] = "user_inferred_value_contents"; +const char kUserInferredValueContentsAttrName[] = + "_user_inferred_value_contents"; const char kUserInferredValueContentsSerializedAttrName[] = + "_user_inferred_value_contents_serialized"; +const char kLegacyUserInferredValueContentsAttrName[] = + "user_inferred_value_contents"; +const char kLegacyUserInferredValueContentsSerializedAttrName[] = "user_inferred_value_contents_serialized"; bool IsShapeOp(const Node* n); @@ -83,6 +88,72 @@ bool GetShapeFromDirectDynamicSource(const Node* node, GetShapeFromArgNode(node, out_shape); } +bool TryParseSerializedContentsAttr(const AttrSlice& attrs, + TensorShapeProto* out_contents) { + string serialized_contents; + if (!GetNodeAttr(attrs, kUserInferredValueContentsSerializedAttrName, + &serialized_contents) + .ok() && + !GetNodeAttr(attrs, kLegacyUserInferredValueContentsSerializedAttrName, + &serialized_contents) + .ok()) { + return false; + } + out_contents->Clear(); + return out_contents->ParseFromString(serialized_contents); +} + +bool TryGetContentsProtoAttr(const AttrSlice& attrs, + TensorShapeProto* out_contents) { + if (TryParseSerializedContentsAttr(attrs, out_contents)) { + return true; + } + if (GetNodeAttr(attrs, kUserInferredValueContentsAttrName, out_contents).ok()) { + return true; + } + return GetNodeAttr(attrs, kLegacyUserInferredValueContentsAttrName, + out_contents) + .ok(); +} + +bool HasTransitiveDynamicShapeContents( + const Node* node, std::unordered_map* memo) { + auto it = memo->find(node); + if (it != memo->end()) { + return it->second; + } + + TensorShapeProto contents_proto; + if (TryGetContentsProtoAttr(node->attrs(), &contents_proto) && + HasDynamicDimExprs(contents_proto)) { + return (*memo)[node] = true; + } + + bool has_dynamic = false; + TensorShapeProto inferred_shape_proto; + if (GetNodeAttr(node->attrs(), "has_dynamic", &has_dynamic).ok() && + has_dynamic && + GetNodeAttr(node->attrs(), "user_inferred_shape", &inferred_shape_proto) + .ok() && + HasDynamicDimExprs(inferred_shape_proto)) { + return (*memo)[node] = true; + } + + if (GetShapeFromDirectDynamicSource(node, &inferred_shape_proto) || + node->attrs().FindByString(kXlaShapeDerivedAttrName) != nullptr) { + return (*memo)[node] = true; + } + + for (const Edge* edge : node->in_edges()) { + if (edge->IsControlEdge()) continue; + if (HasTransitiveDynamicShapeContents(edge->src(), memo)) { + return (*memo)[node] = true; + } + } + + return (*memo)[node] = false; +} + bool GetConstTensor(const Node* node, Tensor* tensor) { if (node == nullptr || !node->IsConstant()) { return false; @@ -192,18 +263,12 @@ bool TryGetFoldedValueContents(const Node* node, int output_index, return false; } - string serialized_contents; - if (GetNodeAttr(node->attrs(), kUserInferredValueContentsSerializedAttrName, - &serialized_contents) - .ok() && - out_contents->ParseFromString(serialized_contents)) { + if (TryParseSerializedContentsAttr(node->attrs(), out_contents)) { return true; } TensorShapeProto existing_contents; - if (GetNodeAttr(node->attrs(), kUserInferredValueContentsAttrName, - &existing_contents) - .ok()) { + if (TryGetContentsProtoAttr(node->attrs(), &existing_contents)) { *out_contents = existing_contents; return true; } @@ -643,11 +708,15 @@ bool IsConstantFoldable( std::unordered_map>* shape_replacement_map) { TensorShapeProto exact_contents; const bool has_exact_contents = TryGetFoldedValueContents(n, 0, &exact_contents); - const bool has_dynamic = - GetShapeFromDirectDynamicSource(n, &exact_contents); + TensorShapeProto dynamic_shape; + const bool has_dynamic = GetShapeFromDirectDynamicSource(n, &dynamic_shape); const bool is_shape_derived = n->attrs().FindByString(kXlaShapeDerivedAttrName) != nullptr; - if ((has_dynamic || is_shape_derived) && (!has_exact_contents || n->num_outputs() > 1)) { + std::unordered_map dynamic_contents_memo; + const bool has_transitive_dynamic_contents = + HasTransitiveDynamicShapeContents(n, &dynamic_contents_memo); + if ((has_dynamic || is_shape_derived || has_transitive_dynamic_contents) && + (!has_exact_contents || n->num_outputs() > 1)) { return false; } if (n->IsConstant()) { diff --git a/tensorflow/core/common_runtime/constant_folding_test.cc b/tensorflow/core/common_runtime/constant_folding_test.cc index 84b5f324f9a62c..7ccde4e6f0b4d4 100644 --- a/tensorflow/core/common_runtime/constant_folding_test.cc +++ b/tensorflow/core/common_runtime/constant_folding_test.cc @@ -680,7 +680,7 @@ TEST_F(ConstantFoldingTest, FoldShapeFromDynamicArgPreservesContents) { // Shape(arg), the replacement Const still carries symbolic contents. string serialized_contents_proto; TF_ASSERT_OK(GetNodeAttr(folded->attrs(), - "user_inferred_value_contents_serialized", + "_user_inferred_value_contents_serialized", &serialized_contents_proto)); TensorShapeProto contents_proto; ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); @@ -729,7 +729,7 @@ TEST_F(ConstantFoldingTest, FoldSliceOfDynamicShapePreservesContents) { // content, not just the concrete value 977. string serialized_contents_proto; TF_ASSERT_OK(GetNodeAttr(folded->attrs(), - "user_inferred_value_contents_serialized", + "_user_inferred_value_contents_serialized", &serialized_contents_proto)); TensorShapeProto contents_proto; ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); @@ -776,7 +776,7 @@ TEST_F(ConstantFoldingTest, FoldGatherOfDynamicShapePreservesContents) { // symbolic content, not just the concrete scalar value 977. string serialized_contents_proto; TF_ASSERT_OK(GetNodeAttr(folded->attrs(), - "user_inferred_value_contents_serialized", + "_user_inferred_value_contents_serialized", &serialized_contents_proto)); TensorShapeProto contents_proto; ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); @@ -826,7 +826,7 @@ TEST_F(ConstantFoldingTest, FoldReshapeOfDynamicShapePreservesContents) { // selected symbolic content when the shape-vector result is scalarized. string serialized_contents_proto; TF_ASSERT_OK(GetNodeAttr(folded->attrs(), - "user_inferred_value_contents_serialized", + "_user_inferred_value_contents_serialized", &serialized_contents_proto)); TensorShapeProto contents_proto; ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); From 07bb8e8d3c6dacaaa6c7f7e96c4dbd7cb57b0816 Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Mon, 27 Apr 2026 13:42:37 +0100 Subject: [PATCH 12/15] Simplify symbolic contents attr contract --- tensorflow/compiler/jit/kernels/xla_ops.cc | 23 ++--------------- .../compiler/tf2xla/kernels/const_op.cc | 23 +---------------- .../core/common_runtime/constant_folding.cc | 25 +++---------------- .../common_runtime/constant_folding_test.cc | 8 +++--- 4 files changed, 11 insertions(+), 68 deletions(-) diff --git a/tensorflow/compiler/jit/kernels/xla_ops.cc b/tensorflow/compiler/jit/kernels/xla_ops.cc index f60f42c3474de8..b828bfab9d75da 100644 --- a/tensorflow/compiler/jit/kernels/xla_ops.cc +++ b/tensorflow/compiler/jit/kernels/xla_ops.cc @@ -103,12 +103,6 @@ namespace tensorflow { namespace { constexpr char kUserInferredValueContentsAttrName[] = "_user_inferred_value_contents"; -constexpr char kUserInferredValueContentsSerializedAttrName[] = - "_user_inferred_value_contents_serialized"; -constexpr char kLegacyUserInferredValueContentsAttrName[] = - "user_inferred_value_contents"; -constexpr char kLegacyUserInferredValueContentsSerializedAttrName[] = - "user_inferred_value_contents_serialized"; using XlaDeviceCompiler = DeviceCompiler; using PjRtDeviceCompiler = @@ -564,16 +558,6 @@ absl::Status CompileToLocalExecutable( auto inferred_shape_it = attr_map.find("user_inferred_shape"); auto inferred_contents_it = attr_map.find(kUserInferredValueContentsAttrName); - if (inferred_contents_it == attr_map.end()) { - inferred_contents_it = - attr_map.find(kLegacyUserInferredValueContentsAttrName); - } - auto inferred_contents_serialized_it = - attr_map.find(kUserInferredValueContentsSerializedAttrName); - if (inferred_contents_serialized_it == attr_map.end()) { - inferred_contents_serialized_it = - attr_map.find(kLegacyUserInferredValueContentsSerializedAttrName); - } bool has_dynamic = false; auto has_dynamic_it = attr_map.find("has_dynamic"); if (has_dynamic_it != attr_map.end()) { @@ -581,19 +565,16 @@ absl::Status CompileToLocalExecutable( } if (inferred_contents_it == attr_map.end() && - inferred_contents_serialized_it == attr_map.end() && (!has_dynamic || inferred_shape_it == attr_map.end())) { return; } TensorShapeProto inferred_shape_proto; - if (inferred_contents_serialized_it != attr_map.end()) { + if (inferred_contents_it != attr_map.end()) { if (!inferred_shape_proto.ParseFromString( - inferred_contents_serialized_it->second.s())) { + inferred_contents_it->second.s())) { return; } - } else if (inferred_contents_it != attr_map.end()) { - inferred_shape_proto = inferred_contents_it->second.shape(); } else { inferred_shape_proto = inferred_shape_it->second.shape(); } diff --git a/tensorflow/compiler/tf2xla/kernels/const_op.cc b/tensorflow/compiler/tf2xla/kernels/const_op.cc index 254f2c26dcb374..2d589302dd62c2 100644 --- a/tensorflow/compiler/tf2xla/kernels/const_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/const_op.cc @@ -37,12 +37,6 @@ namespace { constexpr char kUserInferredValueContentsAttrName[] = "_user_inferred_value_contents"; -constexpr char kUserInferredValueContentsSerializedAttrName[] = - "_user_inferred_value_contents_serialized"; -constexpr char kLegacyUserInferredValueContentsAttrName[] = - "user_inferred_value_contents"; -constexpr char kLegacyUserInferredValueContentsSerializedAttrName[] = - "user_inferred_value_contents_serialized"; template DstT CastTo(SrcT src) { @@ -269,28 +263,13 @@ class ConstOp : public XlaOpKernel { &inferred_shape_proto) .ok()) { } - if (GetNodeAttr(ctx->op_kernel().def(), - kUserInferredValueContentsSerializedAttrName, - &inferred_value_contents_serialized) - .ok() || - GetNodeAttr(ctx->op_kernel().def(), - kLegacyUserInferredValueContentsSerializedAttrName, + if (GetNodeAttr(ctx->op_kernel().def(), kUserInferredValueContentsAttrName, &inferred_value_contents_serialized) .ok()) { if (!inferred_value_contents_proto.ParseFromString( inferred_value_contents_serialized)) { inferred_value_contents_proto.Clear(); } - } else { - GetNodeAttr(ctx->op_kernel().def(), kUserInferredValueContentsAttrName, - &inferred_value_contents_proto) - .IgnoreError(); - if (inferred_value_contents_proto.dim_size() == 0) { - GetNodeAttr(ctx->op_kernel().def(), - kLegacyUserInferredValueContentsAttrName, - &inferred_value_contents_proto) - .IgnoreError(); - } } const bool has_contents_proto = inferred_value_contents_proto.dim_size() > 0; const TensorShapeProto& contents_proto = diff --git a/tensorflow/core/common_runtime/constant_folding.cc b/tensorflow/core/common_runtime/constant_folding.cc index 4bbcd04db68863..ebb050cec35bb3 100644 --- a/tensorflow/core/common_runtime/constant_folding.cc +++ b/tensorflow/core/common_runtime/constant_folding.cc @@ -55,12 +55,6 @@ const char kScopedAllocatorAttrName[] = "_scoped_allocator"; const char kXlaShapeDerivedAttrName[] = "_xla_shape_derived"; const char kUserInferredValueContentsAttrName[] = "_user_inferred_value_contents"; -const char kUserInferredValueContentsSerializedAttrName[] = - "_user_inferred_value_contents_serialized"; -const char kLegacyUserInferredValueContentsAttrName[] = - "user_inferred_value_contents"; -const char kLegacyUserInferredValueContentsSerializedAttrName[] = - "user_inferred_value_contents_serialized"; bool IsShapeOp(const Node* n); @@ -91,10 +85,7 @@ bool GetShapeFromDirectDynamicSource(const Node* node, bool TryParseSerializedContentsAttr(const AttrSlice& attrs, TensorShapeProto* out_contents) { string serialized_contents; - if (!GetNodeAttr(attrs, kUserInferredValueContentsSerializedAttrName, - &serialized_contents) - .ok() && - !GetNodeAttr(attrs, kLegacyUserInferredValueContentsSerializedAttrName, + if (!GetNodeAttr(attrs, kUserInferredValueContentsAttrName, &serialized_contents) .ok()) { return false; @@ -105,15 +96,7 @@ bool TryParseSerializedContentsAttr(const AttrSlice& attrs, bool TryGetContentsProtoAttr(const AttrSlice& attrs, TensorShapeProto* out_contents) { - if (TryParseSerializedContentsAttr(attrs, out_contents)) { - return true; - } - if (GetNodeAttr(attrs, kUserInferredValueContentsAttrName, out_contents).ok()) { - return true; - } - return GetNodeAttr(attrs, kLegacyUserInferredValueContentsAttrName, - out_contents) - .ok(); + return TryParseSerializedContentsAttr(attrs, out_contents); } bool HasTransitiveDynamicShapeContents( @@ -942,7 +925,7 @@ void AddShapeNodeToConstantGraph( .Attr("user_inferred_shape", user_inferred_shape); } if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { - builder.Attr(kUserInferredValueContentsSerializedAttrName, + builder.Attr(kUserInferredValueContentsAttrName, exact_contents.SerializeAsString()); } NodeDef def; @@ -1100,7 +1083,7 @@ bool ReplaceTensorWithConstant( } if (has_exact_contents && HasDynamicDimExprs(exact_contents)) { builder.Attr("has_dynamic", true) - .Attr(kUserInferredValueContentsSerializedAttrName, + .Attr(kUserInferredValueContentsAttrName, exact_contents.SerializeAsString()); } if (partition_device) { diff --git a/tensorflow/core/common_runtime/constant_folding_test.cc b/tensorflow/core/common_runtime/constant_folding_test.cc index 7ccde4e6f0b4d4..18ae0e74a007b6 100644 --- a/tensorflow/core/common_runtime/constant_folding_test.cc +++ b/tensorflow/core/common_runtime/constant_folding_test.cc @@ -680,7 +680,7 @@ TEST_F(ConstantFoldingTest, FoldShapeFromDynamicArgPreservesContents) { // Shape(arg), the replacement Const still carries symbolic contents. string serialized_contents_proto; TF_ASSERT_OK(GetNodeAttr(folded->attrs(), - "_user_inferred_value_contents_serialized", + "_user_inferred_value_contents", &serialized_contents_proto)); TensorShapeProto contents_proto; ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); @@ -729,7 +729,7 @@ TEST_F(ConstantFoldingTest, FoldSliceOfDynamicShapePreservesContents) { // content, not just the concrete value 977. string serialized_contents_proto; TF_ASSERT_OK(GetNodeAttr(folded->attrs(), - "_user_inferred_value_contents_serialized", + "_user_inferred_value_contents", &serialized_contents_proto)); TensorShapeProto contents_proto; ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); @@ -776,7 +776,7 @@ TEST_F(ConstantFoldingTest, FoldGatherOfDynamicShapePreservesContents) { // symbolic content, not just the concrete scalar value 977. string serialized_contents_proto; TF_ASSERT_OK(GetNodeAttr(folded->attrs(), - "_user_inferred_value_contents_serialized", + "_user_inferred_value_contents", &serialized_contents_proto)); TensorShapeProto contents_proto; ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); @@ -826,7 +826,7 @@ TEST_F(ConstantFoldingTest, FoldReshapeOfDynamicShapePreservesContents) { // selected symbolic content when the shape-vector result is scalarized. string serialized_contents_proto; TF_ASSERT_OK(GetNodeAttr(folded->attrs(), - "_user_inferred_value_contents_serialized", + "_user_inferred_value_contents", &serialized_contents_proto)); TensorShapeProto contents_proto; ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); From 1a38be5a0b6c7aad0b1ed787ffeb6571d79ec59b Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Wed, 29 Jul 2026 13:55:18 +0100 Subject: [PATCH 13/15] Clean up symbolic constant folding --- .../compiler/tf2xla/kernels/const_op.cc | 18 ++----- .../core/common_runtime/constant_folding.cc | 53 +++++++++++-------- 2 files changed, 35 insertions(+), 36 deletions(-) diff --git a/tensorflow/compiler/tf2xla/kernels/const_op.cc b/tensorflow/compiler/tf2xla/kernels/const_op.cc index 2d589302dd62c2..a20706a86daade 100644 --- a/tensorflow/compiler/tf2xla/kernels/const_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/const_op.cc @@ -219,16 +219,6 @@ std::vector BuildShapeContentsFromTensorShapeProto( return contents; } -int64_t CountDynamicShapeContents(const TensorShapeProto& shape) { - int64_t dynamic_count = 0; - for (int i = 0; i < shape.expressions_size(); ++i) { - if (IsDynamicExpressionProto(shape.expressions(i))) { - ++dynamic_count; - } - } - return dynamic_count; -} - bool CanAttachContentsFromTensorShapeProto(const TensorShape& tensor_shape, const TensorShapeProto& contents) { return (tensor_shape.dims() == 0 && contents.dim_size() == 1) || @@ -258,10 +248,10 @@ class ConstOp : public XlaOpKernel { TensorShapeProto inferred_value_contents_proto; string inferred_value_contents_serialized; if (GetNodeAttr(ctx->op_kernel().def(), "has_dynamic", &has_dynamic).ok() && - has_dynamic && - GetNodeAttr(ctx->op_kernel().def(), "user_inferred_shape", - &inferred_shape_proto) - .ok()) { + has_dynamic) { + GetNodeAttr(ctx->op_kernel().def(), "user_inferred_shape", + &inferred_shape_proto) + .IgnoreError(); } if (GetNodeAttr(ctx->op_kernel().def(), kUserInferredValueContentsAttrName, &inferred_value_contents_serialized) diff --git a/tensorflow/core/common_runtime/constant_folding.cc b/tensorflow/core/common_runtime/constant_folding.cc index ebb050cec35bb3..68c6bf8df537ea 100644 --- a/tensorflow/core/common_runtime/constant_folding.cc +++ b/tensorflow/core/common_runtime/constant_folding.cc @@ -82,8 +82,8 @@ bool GetShapeFromDirectDynamicSource(const Node* node, GetShapeFromArgNode(node, out_shape); } -bool TryParseSerializedContentsAttr(const AttrSlice& attrs, - TensorShapeProto* out_contents) { +bool TryGetContentsProtoAttr(const AttrSlice& attrs, + TensorShapeProto* out_contents) { string serialized_contents; if (!GetNodeAttr(attrs, kUserInferredValueContentsAttrName, &serialized_contents) @@ -94,17 +94,16 @@ bool TryParseSerializedContentsAttr(const AttrSlice& attrs, return out_contents->ParseFromString(serialized_contents); } -bool TryGetContentsProtoAttr(const AttrSlice& attrs, - TensorShapeProto* out_contents) { - return TryParseSerializedContentsAttr(attrs, out_contents); -} - bool HasTransitiveDynamicShapeContents( - const Node* node, std::unordered_map* memo) { + const Node* node, std::unordered_map* memo, + absl::flat_hash_set* visiting) { auto it = memo->find(node); if (it != memo->end()) { return it->second; } + if (!visiting->insert(node).second) { + return false; + } TensorShapeProto contents_proto; if (TryGetContentsProtoAttr(node->attrs(), &contents_proto) && @@ -129,7 +128,7 @@ bool HasTransitiveDynamicShapeContents( for (const Edge* edge : node->in_edges()) { if (edge->IsControlEdge()) continue; - if (HasTransitiveDynamicShapeContents(edge->src(), memo)) { + if (HasTransitiveDynamicShapeContents(edge->src(), memo, visiting)) { return (*memo)[node] = true; } } @@ -231,7 +230,8 @@ ExpressionProto GetContentExpressionProto(const TensorShapeProto& contents, return MakeConstantExpressionProto(contents.dim(index).size()); } -ExpressionProto MakeMulExpressionProto(ExpressionProto lhs, ExpressionProto rhs) { +ExpressionProto MakeMulExpressionProto(ExpressionProto lhs, + ExpressionProto rhs) { ExpressionProto expr; auto* mul = expr.mutable_mul_node(); *mul->mutable_lhs() = std::move(lhs); @@ -240,19 +240,17 @@ ExpressionProto MakeMulExpressionProto(ExpressionProto lhs, ExpressionProto rhs) } bool TryGetFoldedValueContents(const Node* node, int output_index, - TensorShapeProto* out_contents) { + TensorShapeProto* out_contents, + absl::flat_hash_set* visiting) { out_contents->Clear(); if (output_index != 0) { return false; } - - if (TryParseSerializedContentsAttr(node->attrs(), out_contents)) { - return true; + if (!visiting->insert(node).second) { + return false; } - TensorShapeProto existing_contents; - if (TryGetContentsProtoAttr(node->attrs(), &existing_contents)) { - *out_contents = existing_contents; + if (TryGetContentsProtoAttr(node->attrs(), out_contents)) { return true; } @@ -277,7 +275,7 @@ bool TryGetFoldedValueContents(const Node* node, int output_index, return false; } return TryGetFoldedValueContents(input_edge->src(), input_edge->src_output(), - input_contents); + input_contents, visiting); }; if (node->IsIdentity() || node->type_string() == "Cast") { @@ -328,7 +326,8 @@ bool TryGetFoldedValueContents(const Node* node, int output_index, Tensor axis_tensor; std::vector axis_values; if (!GetInputConstTensor(node, node->num_inputs() - 1, &axis_tensor) || - !GetTensorIntValues(axis_tensor, &axis_values) || axis_values.size() != 1) { + !GetTensorIntValues(axis_tensor, &axis_values) || + axis_values.size() != 1) { return false; } int64_t axis = axis_values[0]; @@ -496,6 +495,12 @@ bool TryGetFoldedValueContents(const Node* node, int output_index, return false; } +bool TryGetFoldedValueContents(const Node* node, int output_index, + TensorShapeProto* out_contents) { + absl::flat_hash_set visiting; + return TryGetFoldedValueContents(node, output_index, out_contents, &visiting); +} + // For stateless RNGs ops, they are pure but device-dependent. Those ops are not // constant-foldable. static absl::flat_hash_set* kBlockList = @@ -690,14 +695,17 @@ bool IsConstantFoldable( int64_t max_constant_size_in_bytes, std::unordered_map>* shape_replacement_map) { TensorShapeProto exact_contents; - const bool has_exact_contents = TryGetFoldedValueContents(n, 0, &exact_contents); + const bool has_exact_contents = + TryGetFoldedValueContents(n, 0, &exact_contents); TensorShapeProto dynamic_shape; const bool has_dynamic = GetShapeFromDirectDynamicSource(n, &dynamic_shape); const bool is_shape_derived = n->attrs().FindByString(kXlaShapeDerivedAttrName) != nullptr; std::unordered_map dynamic_contents_memo; + absl::flat_hash_set dynamic_contents_visiting; const bool has_transitive_dynamic_contents = - HasTransitiveDynamicShapeContents(n, &dynamic_contents_memo); + HasTransitiveDynamicShapeContents(n, &dynamic_contents_memo, + &dynamic_contents_visiting); if ((has_dynamic || is_shape_derived || has_transitive_dynamic_contents) && (!has_exact_contents || n->num_outputs() > 1)) { return false; @@ -910,7 +918,8 @@ void AddShapeNodeToConstantGraph( const bool has_dynamic = GetShapeFromDirectDynamicSource(n, &user_inferred_shape); TensorShapeProto exact_contents; - const bool has_exact_contents = TryGetFoldedValueContents(n, 0, &exact_contents); + const bool has_exact_contents = + TryGetFoldedValueContents(n, 0, &exact_contents); std::vector& added = (*node_map)[n]; const string& node_name = n->name(); for (const Tensor& t : shape_replacement_map.at(n)) { From 0bd94c3c23226d4f431c4be72b8f61022f9af3df Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Wed, 29 Jul 2026 14:33:48 +0100 Subject: [PATCH 14/15] Test symbolic contents through constant folding --- .../common_runtime/constant_folding_test.cc | 82 +++++++++++++++++++ 1 file changed, 82 insertions(+) diff --git a/tensorflow/core/common_runtime/constant_folding_test.cc b/tensorflow/core/common_runtime/constant_folding_test.cc index 18ae0e74a007b6..99d74d6239c1fb 100644 --- a/tensorflow/core/common_runtime/constant_folding_test.cc +++ b/tensorflow/core/common_runtime/constant_folding_test.cc @@ -835,6 +835,88 @@ TEST_F(ConstantFoldingTest, FoldReshapeOfDynamicShapePreservesContents) { EXPECT_EQ(contents_proto.dim(0).size(), 977); } +TEST_F(ConstantFoldingTest, FoldPackOfDynamicShapePreservesContents) { + Graph g(OpRegistry::Global()); + Scope s = Scope::NewRootScope(); + auto arg = ops::_Arg(s.WithOpName("arg"), DT_FLOAT, 0); + auto shape = ops::Shape(s.WithOpName("shape"), arg); + auto first = ops::GatherV2(s.WithOpName("first"), shape, + ops::Const(s.WithOpName("index"), 0), + ops::Const(s.WithOpName("axis"), 0)); + auto second = ops::Const(s.WithOpName("second"), 16); + auto pack = ops::Pack(s.WithOpName("pack"), {first, second}, + ops::Pack::Axis(0)); + auto send = ops::_Send(s.WithOpName("send"), pack, "send", "sender", 0, + "receiver"); + TF_ASSERT_OK(s.ToGraph(&g)); + + std::unordered_map index_by_name = g.BuildNodeNameIndex(); + Node* arg_node = index_by_name.at("arg"); + arg_node->AddAttr("_output_shapes", + std::vector{MakeDynamicShapeProto977x16()}); + + PartialTensorShape partial_shape({977, 16}); + std::unordered_map> shape_map; + shape_map[arg_node->name()].push_back(partial_shape); + + ConstantFoldingOptions opts; + opts.shape_map = &shape_map; + bool was_mutated = false; + TF_ASSERT_OK( + ConstantFold(opts, nullptr, Env::Default(), nullptr, &g, &was_mutated)); + + index_by_name = g.BuildNodeNameIndex(); + Node* send_node = index_by_name.at("send"); + const Edge* send_input = nullptr; + TF_ASSERT_OK(send_node->input_edge(0, &send_input)); + Node* folded = send_input->src(); + ExpectNodeEqual(folded, {977, 16}, {2}); + + string serialized_contents_proto; + TF_ASSERT_OK(GetNodeAttr(folded->attrs(), + "_user_inferred_value_contents", + &serialized_contents_proto)); + TensorShapeProto contents_proto; + ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); + ASSERT_EQ(contents_proto.expressions_size(), 2); + EXPECT_TRUE(IsDynamicDimExpr(contents_proto.expressions(0))); + EXPECT_FALSE(IsDynamicDimExpr(contents_proto.expressions(1))); + EXPECT_EQ(contents_proto.dim(0).size(), 977); + EXPECT_EQ(contents_proto.dim(1).size(), 16); +} + +TEST_F(ConstantFoldingTest, DoNotFoldUnsupportedDynamicContentsTransform) { + Graph g(OpRegistry::Global()); + Scope s = Scope::NewRootScope(); + auto arg = ops::_Arg(s.WithOpName("arg"), DT_FLOAT, 0); + auto shape = ops::Shape(s.WithOpName("shape"), arg); + auto added = ops::Add(s.WithOpName("added"), shape, shape); + auto send = ops::_Send(s.WithOpName("send"), added, "send", "sender", 0, + "receiver"); + TF_ASSERT_OK(s.ToGraph(&g)); + + std::unordered_map index_by_name = g.BuildNodeNameIndex(); + Node* arg_node = index_by_name.at("arg"); + arg_node->AddAttr("_output_shapes", + std::vector{MakeDynamicShapeProto977x16()}); + + PartialTensorShape partial_shape({977, 16}); + std::unordered_map> shape_map; + shape_map[arg_node->name()].push_back(partial_shape); + + ConstantFoldingOptions opts; + opts.shape_map = &shape_map; + bool was_mutated = false; + TF_ASSERT_OK( + ConstantFold(opts, nullptr, Env::Default(), nullptr, &g, &was_mutated)); + + index_by_name = g.BuildNodeNameIndex(); + Node* send_node = index_by_name.at("send"); + const Edge* send_input = nullptr; + TF_ASSERT_OK(send_node->input_edge(0, &send_input)); + EXPECT_EQ(send_input->src()->type_string(), "AddV2"); +} + TEST_F(ConstantFoldingTest, NoReplacePartialOutput) { Graph g(OpRegistry::Global()); { From 129f751ce7e282a27e7b662001268836ff18be3e Mon Sep 17 00:00:00 2001 From: Steven Varoumas Date: Wed, 29 Jul 2026 16:00:46 +0100 Subject: [PATCH 15/15] Fix symbolic constant folding tests --- tensorflow/core/common_runtime/BUILD | 1 + .../core/common_runtime/constant_folding_test.cc | 11 ++++++----- 2 files changed, 7 insertions(+), 5 deletions(-) diff --git a/tensorflow/core/common_runtime/BUILD b/tensorflow/core/common_runtime/BUILD index 24089b11472b6b..f99259ad3e3862 100644 --- a/tensorflow/core/common_runtime/BUILD +++ b/tensorflow/core/common_runtime/BUILD @@ -2774,6 +2774,7 @@ tf_cc_test( "//tensorflow/core/kernels:identity_op", "//tensorflow/core/kernels:immutable_constant_op", "//tensorflow/core/kernels:matmul_op", + "//tensorflow/core/kernels:pack_op", "//tensorflow/core/kernels:reshape_op", "//tensorflow/core/kernels:slice_op", "//tensorflow/core/kernels:topk_op", diff --git a/tensorflow/core/common_runtime/constant_folding_test.cc b/tensorflow/core/common_runtime/constant_folding_test.cc index 99d74d6239c1fb..32014e25615092 100644 --- a/tensorflow/core/common_runtime/constant_folding_test.cc +++ b/tensorflow/core/common_runtime/constant_folding_test.cc @@ -844,8 +844,9 @@ TEST_F(ConstantFoldingTest, FoldPackOfDynamicShapePreservesContents) { ops::Const(s.WithOpName("index"), 0), ops::Const(s.WithOpName("axis"), 0)); auto second = ops::Const(s.WithOpName("second"), 16); - auto pack = ops::Pack(s.WithOpName("pack"), {first, second}, - ops::Pack::Axis(0)); + OutputList pack_inputs = {first, second}; + auto pack = ops::Stack(s.WithOpName("pack"), pack_inputs, + ops::Stack::Axis(0)); auto send = ops::_Send(s.WithOpName("send"), pack, "send", "sender", 0, "receiver"); TF_ASSERT_OK(s.ToGraph(&g)); @@ -870,6 +871,7 @@ TEST_F(ConstantFoldingTest, FoldPackOfDynamicShapePreservesContents) { const Edge* send_input = nullptr; TF_ASSERT_OK(send_node->input_edge(0, &send_input)); Node* folded = send_input->src(); + ASSERT_TRUE(folded->IsConstant()); ExpectNodeEqual(folded, {977, 16}, {2}); string serialized_contents_proto; @@ -878,9 +880,8 @@ TEST_F(ConstantFoldingTest, FoldPackOfDynamicShapePreservesContents) { &serialized_contents_proto)); TensorShapeProto contents_proto; ASSERT_TRUE(contents_proto.ParseFromString(serialized_contents_proto)); - ASSERT_EQ(contents_proto.expressions_size(), 2); + ASSERT_EQ(contents_proto.expressions_size(), 1); EXPECT_TRUE(IsDynamicDimExpr(contents_proto.expressions(0))); - EXPECT_FALSE(IsDynamicDimExpr(contents_proto.expressions(1))); EXPECT_EQ(contents_proto.dim(0).size(), 977); EXPECT_EQ(contents_proto.dim(1).size(), 16); } @@ -914,7 +915,7 @@ TEST_F(ConstantFoldingTest, DoNotFoldUnsupportedDynamicContentsTransform) { Node* send_node = index_by_name.at("send"); const Edge* send_input = nullptr; TF_ASSERT_OK(send_node->input_edge(0, &send_input)); - EXPECT_EQ(send_input->src()->type_string(), "AddV2"); + EXPECT_FALSE(send_input->src()->IsConstant()); } TEST_F(ConstantFoldingTest, NoReplacePartialOutput) {