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/BUILD b/tensorflow/compiler/jit/kernels/BUILD index 8da9744bacd79e..1afe7ccd2fb90e 100644 --- a/tensorflow/compiler/jit/kernels/BUILD +++ b/tensorflow/compiler/jit/kernels/BUILD @@ -1,3 +1,4 @@ +load("//tensorflow:tensorflow.bzl", "tf_cc_test") load("//tensorflow/core/platform:rules_cc.bzl", "cc_library") package( @@ -48,7 +49,10 @@ XLA_OPS_DEPS = [ cc_library( name = "xla_ops_no_jit_rewrite_registration", srcs = ["xla_ops.cc"], - hdrs = ["xla_ops.h"], + hdrs = [ + "xla_ops.h", + "xla_ops_internal.h", + ], deps = XLA_OPS_DEPS + [ "//tensorflow/compiler/jit:device_compilation_cache", "//tensorflow/compiler/jit:device_compilation_profiler", @@ -75,6 +79,20 @@ cc_library( alwayslink = 1, ) +tf_cc_test( + name = "xla_ops_test", + srcs = ["xla_ops_test.cc"], + env = {"TF_XLA_FLAGS": "--tf_xla_enable_dynamic_sizes"}, + deps = [ + ":xla_ops_no_jit_rewrite_registration", + "//tensorflow/compiler/tf2xla:xla_argument", + "//tensorflow/core:framework", + "//tensorflow/core/platform:test", + "@com_google_googletest//:gtest_main", + "@local_xla//xla:shape_util", + ], +) + cc_library( name = "xla_ops", hdrs = ["xla_ops.h"], diff --git a/tensorflow/compiler/jit/kernels/xla_ops.cc b/tensorflow/compiler/jit/kernels/xla_ops.cc index 2fb97d66b3ac63..95a3a6d4f84844 100644 --- a/tensorflow/compiler/jit/kernels/xla_ops.cc +++ b/tensorflow/compiler/jit/kernels/xla_ops.cc @@ -14,6 +14,7 @@ limitations under the License. ==============================================================================*/ #include "tensorflow/compiler/jit/kernels/xla_ops.h" +#include "tensorflow/compiler/jit/kernels/xla_ops_internal.h" #include #include @@ -36,6 +37,7 @@ limitations under the License. #include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/str_cat.h" +#include "absl/strings/str_join.h" #include "absl/strings/string_view.h" #include "absl/synchronization/mutex.h" #include "absl/types/span.h" @@ -101,6 +103,8 @@ limitations under the License. namespace tensorflow { namespace { +constexpr char kUserInferredValueContentsAttrName[] = + "_user_inferred_value_contents"; using XlaDeviceCompiler = DeviceCompiler; using PjRtDeviceCompiler = @@ -205,6 +209,571 @@ using PjRtExecutableClosure = using PjRtExecutableClosureStore = ExecutableClosureStore; +struct DynamicBatchResolutionResult { + bool can_run = true; + bool has_batch_size = false; + int64_t batch_size = 0; + bool batch_size_resource_not_found = false; + std::string diagnostic; +}; + +struct DynamicSolveEvidence { + int64_t solved_value; + int64_t observed_value; + std::string expr; +}; + +struct DynamicSolveSelection { + std::set singleton_values; + std::set nonsingleton_values; + std::set expressions; + std::optional chosen_value; + std::string reason; + bool rejected = false; +}; + +template +DynamicSolveSelection SelectDynamicSolveEvidence( + absl::Span candidates) { + DynamicSolveSelection selection; + for (const Candidate& candidate : candidates) { + const DynamicSolveEvidence& evidence = candidate.evidence; + selection.expressions.insert(evidence.expr); + if (evidence.observed_value == 1) { + selection.singleton_values.insert(evidence.solved_value); + } else { + selection.nonsingleton_values.insert(evidence.solved_value); + } + } + + // A runtime singleton may be an implicitly broadcast operand. Prefer + // consistent non-singleton evidence for the same normalized variable. + if (!selection.nonsingleton_values.empty()) { + if (selection.nonsingleton_values.size() != 1) { + selection.rejected = true; + selection.reason = "conflicting non-singleton candidates"; + return selection; + } + selection.chosen_value = *selection.nonsingleton_values.begin(); + selection.reason = + selection.singleton_values.empty() + ? "used non-singleton evidence" + : "preferred non-singleton evidence over singleton-derived " + "candidates"; + return selection; + } + + if (selection.singleton_values.size() == 1) { + selection.chosen_value = *selection.singleton_values.begin(); + selection.reason = "used singleton-derived evidence only"; + } else if (selection.singleton_values.size() > 1) { + selection.rejected = true; + selection.reason = "conflicting singleton-derived candidates"; + } else { + selection.reason = "no dynamic values were solved"; + } + return selection; +} + +struct DynamicSolveCandidate { + DynamicSolveEvidence evidence; + std::string context; +}; + +} // namespace + +namespace xla_ops_internal { + +int GetRuntimeInputIndex(absl::Span input_mapping, + int xla_input_index, int num_constant_args, + bool constants_omitted) { + if (xla_input_index < 0 || xla_input_index >= input_mapping.size()) { + return -1; + } + const int missing_input_prefix = constants_omitted ? num_constant_args : 0; + return input_mapping[xla_input_index] - missing_input_prefix; +} + +std::optional GetConstantArgumentElementValue( + const XlaArgument& arg, int index) { + if (index < 0 || index >= arg.constant_value.NumElements()) { + return std::nullopt; + } + switch (arg.constant_value.dtype()) { + case DT_INT32: + return static_cast(arg.constant_value.flat()(index)); + case DT_INT64: + return static_cast(arg.constant_value.flat()(index)); + default: + return std::nullopt; + } +} + +void SetConstantArgumentExpressionToLiteralValue(XlaArgument* arg, + int index) { + if (arg == nullptr || index < 0 || + index >= arg->constant_value_expressions.size()) { + return; + } + std::optional value = GetConstantArgumentElementValue(*arg, index); + if (!value.has_value()) { + arg->constant_value_expressions[index].Clear(); + return; + } + arg->constant_value_expressions[index].Clear(); + arg->constant_value_expressions[index].set_constant_value(*value); +} + +DynamicSolveFilterDecision AnalyzeIgnoredDynamicArgumentOccurrences( + absl::Span args) { + struct Candidate { + DynamicSolveEvidence evidence; + IgnoredDynamicArgumentOccurrence occurrence; + }; + + DynamicSolveFilterDecision result; + std::map, std::vector> variable_ids_to_candidates; + + for (int arg_index = 0; arg_index < args.size(); ++arg_index) { + const XlaArgument& arg = args[arg_index]; + if (absl::holds_alternative(arg.shape)) { + const TensorShape& shape = std::get(arg.shape); + for (int dim = 0; dim < shape.get_expressions().size(); ++dim) { + const xla::DExpr& expr = shape.get_expression(dim); + if (!(expr && expr->is_dynamic())) { + continue; + } + xla::DExpr simplified_expr = expr.simplify(); + std::optional solved_value = + simplified_expr->solve(shape.dim_size(dim)); + if (!solved_value.has_value()) { + continue; + } + const std::string expr_string = DExprToString(simplified_expr); + variable_ids_to_candidates[simplified_expr->get_all_ids()].push_back( + Candidate{ + DynamicSolveEvidence{*solved_value, shape.dim_size(dim), + expr_string}, + IgnoredDynamicArgumentOccurrence{ + IgnoredDynamicArgumentOccurrence::Source::kShapeDimension, + arg_index, + dim, + shape.dim_size(dim), + *solved_value, + expr_string}}); + } + } + + if (arg.kind != XlaCompiler::Argument::kConstant || + arg.constant_value_expressions.empty()) { + continue; + } + for (int element_index = 0; + element_index < arg.constant_value_expressions.size(); + ++element_index) { + xla::DExpr expr = + xla::DExprFromProto(arg.constant_value_expressions[element_index]); + if (!(expr && expr->is_dynamic())) { + continue; + } + std::optional observed_value = + GetConstantArgumentElementValue(arg, element_index); + if (!observed_value.has_value()) { + continue; + } + xla::DExpr simplified_expr = expr.simplify(); + std::optional solved_value = + simplified_expr->solve(*observed_value); + if (!solved_value.has_value()) { + continue; + } + const std::string expr_string = DExprToString(simplified_expr); + variable_ids_to_candidates[simplified_expr->get_all_ids()].push_back( + Candidate{ + DynamicSolveEvidence{*solved_value, *observed_value, + expr_string}, + IgnoredDynamicArgumentOccurrence{ + IgnoredDynamicArgumentOccurrence::Source:: + kConstantValueElement, + arg_index, + element_index, + *observed_value, + *solved_value, + expr_string}}); + } + } + + std::vector expr_summaries; + for (const auto& [variable_ids, candidates] : variable_ids_to_candidates) { + DynamicSolveSelection selection = SelectDynamicSolveEvidence( + absl::MakeConstSpan(candidates)); + + if (selection.rejected) { + result.can_run = false; + } else if (selection.chosen_value.has_value() && + !selection.nonsingleton_values.empty()) { + for (const auto& candidate : candidates) { + if (candidate.evidence.observed_value == 1 && + candidate.evidence.solved_value != *selection.chosen_value) { + VLOG(1) << "Ignoring singleton-derived dynamic " + << "occurrence during XLA signature filtering: " + << "variable_ids={" << absl::StrJoin(variable_ids, ", ") + << "} expr=" << candidate.occurrence.expr + << " chosen_value=" << *selection.chosen_value + << " ignored_observed_value=" + << candidate.evidence.observed_value + << " ignored_solved_value=" + << candidate.evidence.solved_value + << " arg_index=" << candidate.occurrence.arg_index + << (candidate.occurrence.source == + IgnoredDynamicArgumentOccurrence::Source:: + kShapeDimension + ? " dim=" + : " element=") + << candidate.occurrence.dim_or_index; + result.ignored_occurrences.push_back(candidate.occurrence); + } + } + } + + expr_summaries.push_back(absl::StrCat( + "variable_ids={", absl::StrJoin(variable_ids, ", "), + "} expressions={", absl::StrJoin(selection.expressions, ", "), + "} chosen=", + selection.chosen_value.has_value() + ? std::to_string(*selection.chosen_value) + : std::string(""), + " singleton_derived_values={", + absl::StrJoin(selection.singleton_values, ", "), + "} nonsingleton_values={", + absl::StrJoin(selection.nonsingleton_values, ", "), "} reason=", + selection.reason)); + } + + if (!result.can_run) { + result.diagnostic = absl::StrCat( + "Failed to recover a unique XLA dynamic batch size after solve-time " + "candidate filtering. solved_expressions=[", + absl::StrJoin(expr_summaries, " | "), "]"); + return result; + } + + if (!result.ignored_occurrences.empty()) { + std::vector ignored_summaries; + ignored_summaries.reserve(result.ignored_occurrences.size()); + for (const auto& occurrence : result.ignored_occurrences) { + ignored_summaries.push_back(absl::StrCat( + occurrence.source == + IgnoredDynamicArgumentOccurrence::Source::kShapeDimension + ? "shape_arg=" + : "const_arg=", + occurrence.arg_index, + occurrence.source == + IgnoredDynamicArgumentOccurrence::Source::kShapeDimension + ? " dim=" + : " element=", + occurrence.dim_or_index, " expr=", occurrence.expr, + " observed=", occurrence.observed_value, + " solved=", occurrence.solved_value)); + } + result.diagnostic = absl::StrCat( + "Ignoring singleton-derived dynamic occurrences for XLA " + "signature/HLO: ", + absl::StrJoin(ignored_summaries, "; ")); + } + return result; +} + +std::vector BuildStaticCompilationArguments( + absl::Span args) { + std::vector static_args(args.begin(), args.end()); + auto clear_shape_expressions = [&](auto&& self, xla::Shape* shape) -> void { + if (shape->IsTuple()) { + for (xla::Shape& subshape : *shape->mutable_tuple_shapes()) { + self(self, &subshape); + } + } else if (shape->IsArray()) { + shape->set_expressions({}); + for (int dim = 0; dim < shape->dimensions().size(); ++dim) { + if (shape->dimensions(dim) >= 0) { + shape->set_dynamic_dimension(dim, false); + } + } + } + }; + + // Preserve concrete runtime dimensions and values so CompileIfNeeded uses + // its ordinary static-shape cache key. Only symbolic/dynamic annotations + // are removed. + for (XlaArgument& arg : static_args) { + if (absl::holds_alternative(arg.shape)) { + std::get(arg.shape).set_expressions({}); + } else { + clear_shape_expressions(clear_shape_expressions, + &std::get(arg.shape)); + } + arg.constant_value_expressions.clear(); + } + return static_args; +} + +void StripIgnoredDynamicArgumentOccurrences( + const std::vector& ignored_occurrences, + std::vector* args) { + if (args == nullptr) { + return; + } + for (const auto& occurrence : ignored_occurrences) { + if (occurrence.arg_index < 0 || occurrence.arg_index >= args->size()) { + continue; + } + XlaArgument& arg = (*args)[occurrence.arg_index]; + if (occurrence.source == + IgnoredDynamicArgumentOccurrence::Source::kShapeDimension) { + if (!absl::holds_alternative(arg.shape)) { + continue; + } + VLOG(1) << "Dropping ignored dynamic shape expression from XLA " + << "signature/HLO: arg_index=" << occurrence.arg_index + << " dim=" << occurrence.dim_or_index + << " expr=" << occurrence.expr + << " observed_value=" << occurrence.observed_value + << " solved_value=" << occurrence.solved_value; + TensorShape& shape = std::get(arg.shape); + shape.set_expression(occurrence.dim_or_index, xla::DExpr()); + continue; + } + + VLOG(1) << "Dropping ignored dynamic constant-value expression from XLA " + << "signature/HLO: arg_index=" << occurrence.arg_index + << " element=" << occurrence.dim_or_index + << " expr=" << occurrence.expr + << " observed_value=" << occurrence.observed_value + << " solved_value=" << occurrence.solved_value; + SetConstantArgumentExpressionToLiteralValue(&arg, occurrence.dim_or_index); + } +} + +} // namespace xla_ops_internal + +namespace { + +using xla_ops_internal::AnalyzeIgnoredDynamicArgumentOccurrences; +using xla_ops_internal::BuildStaticCompilationArguments; +using xla_ops_internal::DynamicSolveFilterDecision; +using xla_ops_internal::GetRuntimeInputIndex; +using xla_ops_internal::StripIgnoredDynamicArgumentOccurrences; + +DynamicBatchResolutionResult ResolveDynamicBatchSizeFromRuntimeInputs( + OpKernelContext* ctx, const XlaCompiler::CompilationResult& comp_result, + int num_constant_args, bool log_solves, absl::string_view cluster_name, + const NodeDef& op_def) { + DynamicBatchResolutionResult result; + MarkForCompilationPassFlags* flags = GetMarkForCompilationPassFlags(); + if (!flags->tf_xla_enable_dynamic_sizes) { + return result; + } + auto resolve_runtime_input_index = [&](int xla_input_index) -> int { + // input_mapping maps each XLA parameter to its original function argument. + // _XlaRun omits the compile-time constant prefix from its inputs. + return GetRuntimeInputIndex(comp_result.input_mapping, xla_input_index, + num_constant_args, + /*constants_omitted=*/op_def.op() == "_XlaRun"); + }; + + std::set dyn_vals; + std::map, std::vector> + variable_ids_to_candidates; + + for (int i = 0; i < comp_result.xla_input_shapes.size(); ++i) { + const auto& xla_shape = comp_result.xla_input_shapes[i]; + const int input_idx = resolve_runtime_input_index(i); + if (!xla_shape.IsArray() || xla_shape.expressions().empty()) { + continue; + } + + for (int dim = 0; dim < xla_shape.expressions().size(); ++dim) { + const auto& expr = xla_shape.expressions(dim); + if (!(expr && expr->is_dynamic())) { + continue; + } + xla::DExpr simplified_expr = expr.simplify(); + if (input_idx < 0 || input_idx >= ctx->num_inputs()) { + continue; + } + const bool is_runtime_key_input = + ctx->input_dtype(input_idx) == DT_STRING && + input_idx == op_def.input_size() - 1; + const std::string runtime_input_name = + input_idx < op_def.input_size() ? op_def.input(input_idx) + : std::string(""); + if (is_runtime_key_input) { + if (log_solves) { + VLOG(1) << "Skipping dynamic expression solve for runtime key input " + << "cluster=" << cluster_name + << " xla_input_index=" << i + << " runtime_input_index=" << input_idx + << " input_name=" << runtime_input_name; + } + continue; + } + int64_t size = ctx->input(input_idx).shape().dim_size(dim); + std::optional dyn_val = simplified_expr->solve(size); + if (log_solves && dyn_val.has_value()) { + VLOG(1) << "Dynamic solve: cluster=" << cluster_name + << " runtime_input_index=" << input_idx + << " xla_input_index=" << i << " dim=" << dim + << " input_name=" << runtime_input_name + << " expr=" << DExprToString(simplified_expr) + << " target_size=" << size << " result=" + << (dyn_val.has_value() ? std::to_string(*dyn_val) + : std::string("")); + } + if (!dyn_val.has_value()) { + if (log_solves) { + LOG(WARNING) << "Dynamic solve failed: cluster=" << cluster_name + << " runtime_input_index=" << input_idx + << " xla_input_index=" << i << " dim=" << dim + << " input_name=" << runtime_input_name + << " expr=" << DExprToString(simplified_expr) + << " target_size=" << size; + } + continue; + } + const std::string expr_string = DExprToString(simplified_expr); + const std::string context = absl::StrCat( + "xla_input_index=", i, " runtime_input_index=", input_idx, + " input_name=", runtime_input_name, " dim=", dim, + " runtime_dim_size=", size, + " solved_dynamic_value=", *dyn_val); + variable_ids_to_candidates[simplified_expr->get_all_ids()].push_back( + DynamicSolveCandidate{ + DynamicSolveEvidence{*dyn_val, size, expr_string}, context}); + } + } + + std::vector expr_summaries; + std::optional> expected_dyn_ids; + bool mismatched_dyn_ids = false; + bool candidate_filter_rejected = false; + for (const auto& [variable_ids, candidates] : variable_ids_to_candidates) { + std::vector singleton_contexts; + std::vector nonsingleton_contexts; + for (const auto& candidate : candidates) { + if (candidate.evidence.observed_value == 1) { + singleton_contexts.push_back(candidate.context); + } else { + nonsingleton_contexts.push_back(candidate.context); + } + } + DynamicSolveSelection selection = SelectDynamicSolveEvidence( + absl::MakeConstSpan(candidates)); + + if (selection.rejected) { + candidate_filter_rejected = true; + LOG(WARNING) << "Dynamic solve rejected candidates for variable_ids={" + << absl::StrJoin(variable_ids, ", ") + << "}: " << selection.reason; + } + + if (selection.chosen_value.has_value()) { + dyn_vals.insert(*selection.chosen_value); + } + + expr_summaries.push_back(absl::StrCat( + "variable_ids={", absl::StrJoin(variable_ids, ", "), + "} expressions={", absl::StrJoin(selection.expressions, ", "), + "} chosen=", + selection.chosen_value.has_value() + ? std::to_string(*selection.chosen_value) + : std::string(""), + " singleton_derived_values={", + absl::StrJoin(selection.singleton_values, ", "), + "} nonsingleton_values={", + absl::StrJoin(selection.nonsingleton_values, ", "), "} reason=", + selection.reason, + " singleton_derived_contexts=[", + absl::StrJoin(singleton_contexts, "; "), "] nonsingleton_contexts=[", + absl::StrJoin(nonsingleton_contexts, "; "), "]")); + } + for (int i = 0; i < comp_result.xla_input_shapes.size(); ++i) { + const auto& xla_shape = comp_result.xla_input_shapes[i]; + if (!xla_shape.IsArray() || xla_shape.expressions().empty()) continue; + for (int dim = 0; dim < xla_shape.expressions().size(); ++dim) { + const auto& expr = xla_shape.expressions(dim); + if (!(expr && expr->is_dynamic())) { + continue; + } + std::set ids = expr->get_all_ids(); + if (!expected_dyn_ids.has_value()) { + expected_dyn_ids = ids; + } else if (*expected_dyn_ids != ids) { + mismatched_dyn_ids = true; + } + } + } + + if (dyn_vals.size() == 1 && expected_dyn_ids.has_value() && + expected_dyn_ids->size() == 1 && !mismatched_dyn_ids) { + result.has_batch_size = true; + result.batch_size = *dyn_vals.begin(); + result.diagnostic = absl::StrCat( + "shared variable ids={", absl::StrJoin(*expected_dyn_ids, ", "), + "} value=", result.batch_size); + return result; + } + + if (variable_ids_to_candidates.empty()) { + result.diagnostic = "no dynamic values were solved"; + } else if (candidate_filter_rejected) { + result.can_run = false; + result.diagnostic = absl::StrCat( + "Failed to recover a unique XLA dynamic batch size after solve-time " + "candidate filtering. solved_expressions=[", + absl::StrJoin(expr_summaries, " | "), "] all_values={", + absl::StrJoin(dyn_vals, ", "), "}"); + return result; + } else if (dyn_vals.size() == 1 && mismatched_dyn_ids) { + result.diagnostic = absl::StrCat( + "solved dynamic expressions do not share the same variable ids: ", + absl::StrJoin(expr_summaries, " | ")); + } else if (dyn_vals.size() == 1) { + result.diagnostic = absl::StrCat( + "solved dynamic expressions do not map to exactly one shared dynamic " + "variable id: ", + absl::StrJoin(expr_summaries, " | ")); + } else { + result.can_run = false; + result.diagnostic = absl::StrCat( + "Failed to recover a unique XLA dynamic batch size from runtime " + "input expressions. solved_expressions=[", + absl::StrJoin(expr_summaries, " | "), "] all_values={", + absl::StrJoin(dyn_vals, ", "), "}"); + return result; + } + + BatchSizeResource* bsr = nullptr; + ScopedStepContainer* step_container = ctx->step_container(); + absl::Status st = step_container->Lookup( + ctx->resource_manager(), BatchSizeResourceName, &bsr); + + if (st.ok()) { + CHECK(bsr != nullptr); + result.has_batch_size = true; + result.batch_size = bsr->GetBatchSize(); + result.diagnostic = absl::StrCat( + "BatchSizeResource fallback value=", result.batch_size); + bsr->Unref(); + } else if (IsNotFound(st)) { + result.batch_size_resource_not_found = true; + } else { + result.can_run = false; + result.diagnostic = st.ToString(); + } + + return result; +} + se::Stream* GetStream(OpKernelContext* ctx) { return ctx->op_device_context() ? ctx->op_device_context()->stream() : nullptr; @@ -385,76 +954,12 @@ GetXlaCompilerArgsAndSnapshotVariables( return result; } - -std::unique_ptr ExprFromProto(const ExpressionProto& proto) { - switch (proto.node_type_case()) { - case ExpressionProto::kConstantValue: - return DimExpr::Cons(proto.constant_value()); - case ExpressionProto::kVariableId: - return DimExpr::Var(proto.variable_id()); - case ExpressionProto::kAddNode: { - auto lhs = ExprFromProto(proto.add_node().lhs()); - auto rhs = ExprFromProto(proto.add_node().rhs()); - // Note: These are owning pointers, but ExprAdd takes raw pointers. - // The caller must manage lifetime appropriately. - return std::make_unique(lhs.release(), rhs.release()); - } - case ExpressionProto::kSubNode: { - auto lhs = ExprFromProto(proto.sub_node().lhs()); - auto rhs = ExprFromProto(proto.sub_node().rhs()); - return std::make_unique(lhs.release(), rhs.release()); - } - case ExpressionProto::kMulNode: { - auto lhs = ExprFromProto(proto.mul_node().lhs()); - auto rhs = ExprFromProto(proto.mul_node().rhs()); - return std::make_unique(lhs.release(), rhs.release()); - } - case ExpressionProto::kDivNode: { - auto lhs = ExprFromProto(proto.div_node().lhs()); - auto rhs = ExprFromProto(proto.div_node().rhs()); - return std::make_unique(lhs.release(), rhs.release()); - } - case ExpressionProto::NODE_TYPE_NOT_SET: - default: - return nullptr; - } -} - -static xla::DExpr DimExprToDExpr(const DimExpr* e) { - switch (e->kind()) { - case DimExpr::Kind::kConstant: { - auto* ac = static_cast(e); - return xla::DExpr::Const(ac->value()); - } - case DimExpr::Kind::kVariable: { - return xla::DExpr::Var(1); - } - case DimExpr::Kind::kAdd: { - auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) + DimExprToDExpr(ee->rhs()); - } - case DimExpr::Kind::kSub: { - auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) - DimExprToDExpr(ee->rhs()); - } - case DimExpr::Kind::kMul: { - auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) * DimExprToDExpr(ee->rhs()); - } - case DimExpr::Kind::kDiv: { - auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) / DimExprToDExpr(ee->rhs()); - } - } - return xla::DExpr::Unknown(); -} - - absl::Status CompileToLocalExecutable( OpKernelContext* ctx, const NameAttrList& function, bool has_ref_vars, const XlaPlatformInfo& platform_info, const std::vector& args, DeviceCompileMode compile_mode, bool may_alias_resource_update, + bool force_static_shapes, bool* dynamic_solve_conflict, xla::LocalClient** client, const XlaCompiler::CompilationResult** compilation_result, xla::LocalExecutable** executable) { @@ -498,8 +1003,11 @@ absl::Status CompileToLocalExecutable( XlaCompiler::CompileOptions compile_options = GenerateCompileOptions(has_ref_vars, may_alias_resource_update); + if (dynamic_solve_conflict != nullptr) { + *dynamic_solve_conflict = false; + } MarkForCompilationPassFlags* flags = GetMarkForCompilationPassFlags(); - if (flags->tf_xla_enable_dynamic_sizes) { + if (flags->tf_xla_enable_dynamic_sizes && !force_static_shapes) { // Rewriting the argument with expressions if they have dynamic // dimension, detecting dynamic dimension via either _dynamic_dim or the // inferred-output-shapes attr attached during encapsulation. @@ -512,6 +1020,31 @@ absl::Status CompileToLocalExecutable( XlaBatchMatcher* xla_batch_matcher = xla_device_compiler->xla_batch_matcher(); std::optional dynamic_dim_expr; + std::optional shared_dynamic_subexpr; + auto normalize_dynamic_expr = + [&](xla::DExpr expr, absl::string_view context = "") { + if (!expr || !expr->is_dynamic()) { + return expr; + } + if (!shared_dynamic_subexpr.has_value()) { + shared_dynamic_subexpr = + expr.find_smallest_subexpression_covering_all_variables(); + VLOG(1) << "Using shared dynamic subexpression " + << DExprToString(*shared_dynamic_subexpr) + << " for XLA dynamic input normalization. context=" + << context; + } + xla::DExpr normalized_expr = + expr.replace_subexpression(*shared_dynamic_subexpr, + xla::DExpr::Var(1)) + .simplify(); + VLOG(1) << "Rewriting dynamic input expression " + << DExprToString(expr) + << " to shared-core form " + << DExprToString(normalized_expr) + << " before XLA compilation. context=" << context; + return normalized_expr; + }; auto maybe_attach_shape_contents_from_attrs = [&](int arg_index, const auto& attr_map, const std::string& node_name) { @@ -520,35 +1053,35 @@ absl::Status CompileToLocalExecutable( return; } + auto inferred_shape_it = attr_map.find("user_inferred_shape"); + auto inferred_contents_it = + attr_map.find(kUserInferredValueContentsAttrName); 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()) { + if (!inferred_shape_proto.ParseFromString( + inferred_contents_it->second.s())) { + return; + } + } 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()) { - 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(); + if (!((TensorShapeUtils::IsVector(arg.constant_value.shape()) && + arg.constant_value.NumElements() == inferred_shape.dims()) || + (TensorShapeUtils::IsScalar(arg.constant_value.shape()) && + inferred_shape.dims() == 1))) { return; } @@ -558,24 +1091,21 @@ absl::Status CompileToLocalExecutable( xla::ExpressionProto expr; const xla::DExpr& dim_expr = inferred_shape.get_expression(i); if (dim_expr && dim_expr->is_dynamic()) { - dim_expr->to_proto(&expr); + xla::DExpr normalized_expr = normalize_dynamic_expr( + dim_expr, + absl::StrCat("const_arg=", arg_index, " node=", node_name, + " dim=", i)); + normalized_expr->to_proto(&expr); } else if (arg.constant_value.dtype() == DT_INT32) { expr.set_constant_value(arg.constant_value.flat()(i)); } 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) { @@ -621,7 +1151,7 @@ absl::Status CompileToLocalExecutable( for (int d : shp.dim_sizes()) { dyn_exprs.push_back(xla::DExpr::Const(d)); } - dyn_exprs[idx] = xla::DExpr::Var(1); + dyn_exprs[idx] = *dynamic_dim_expr; shp.set_expressions(std::move(dyn_exprs)); continue; } @@ -631,15 +1161,23 @@ absl::Status CompileToLocalExecutable( const TensorShapeProto& proto = it->second.list().shape(0); const auto& exp = proto.expressions(); TensorShape& shp = std::get(norm_args[arg_index].shape); - if (!filled_batch && xla_batch_matcher) { for (int idx = 0; idx < exp.size(); ++idx) { // Look for dynamic expression. If found then compute padding // value and exit loop. - auto e = DimExprToDExpr(ExprFromProto(exp[idx]).get()).simplify(); + auto e = normalize_dynamic_expr(DimExprFromProto(exp[idx])); if (e->is_dynamic()) { + VLOG(1) << "Calling dynamic expression solve for compile " + << "argument " << arg_index << " dimension " << idx + << " expr=" << DExprToString(e) + << " target_size=" << shp.dim_size(idx); std::optional solved_value = e->solve(shp.dim_size(idx)); + VLOG(1) << "Dynamic expression solve for compile argument " + << arg_index << " dimension " << idx << " returned " + << (solved_value.has_value() + ? std::to_string(*solved_value) + : std::string("")); int64_t var_value; if (!solved_value.has_value()) { LOG(WARNING) @@ -666,8 +1204,11 @@ absl::Status CompileToLocalExecutable( dyn_exprs.push_back(xla::DExpr::Const(d)); } for (int j = 0; j < exp.size(); ++j) { - auto e = DimExprToDExpr(ExprFromProto(exp[j]).get()); + auto e = DimExprFromProto(exp[j]); if (e->is_dynamic()) { + e = normalize_dynamic_expr( + e, absl::StrCat("arg=", arg_index, " dim=", j, + " input_shape")); dyn_exprs[j] = e; } } @@ -676,6 +1217,51 @@ absl::Status CompileToLocalExecutable( } } + DynamicSolveFilterDecision solve_filter_decision = + AnalyzeIgnoredDynamicArgumentOccurrences(norm_args); + if (!solve_filter_decision.can_run) { + if (dynamic_solve_conflict != nullptr) { + *dynamic_solve_conflict = true; + } + if (compile_mode == DeviceCompileMode::kLazy) { + return errors::Unimplemented(solve_filter_decision.diagnostic); + } + return errors::InvalidArgument(solve_filter_decision.diagnostic); + } + if (!solve_filter_decision.ignored_occurrences.empty()) { + LOG(WARNING) << solve_filter_decision.diagnostic; + StripIgnoredDynamicArgumentOccurrences( + solve_filter_decision.ignored_occurrences, &norm_args); + saw_dynamic_dim_value = false; + has_multiple_dynamic_dim_values = false; + dynamic_dim_value = 0; + dynamic_dim_expr.reset(); + filled_batch = 0; + for (int arg_index = 0; arg_index < norm_args.size(); ++arg_index) { + if (!absl::holds_alternative(norm_args[arg_index].shape)) { + continue; + } + TensorShape& shape = std::get(norm_args[arg_index].shape); + for (int dim = 0; dim < shape.get_expressions().size(); ++dim) { + xla::DExpr expr = shape.get_expression(dim); + if (!(expr && expr->is_dynamic())) { + continue; + } + xla::DExpr simplified_expr = expr.simplify(); + std::optional solved_value = + simplified_expr->solve(shape.dim_size(dim)); + if (!solved_value.has_value()) { + continue; + } + record_dynamic_dim_value(*solved_value, simplified_expr); + if (!filled_batch && xla_batch_matcher) { + filled_batch = + xla_batch_matcher->get_xla_compile_batch(*solved_value); + } + } + } + } + struct SaveOldVar { int arg_index; int64_t dyn_dim; @@ -771,7 +1357,25 @@ absl::Status CompileToLocalExecutable( int64_t old = shp.dim_size(j); old_vars.push_back({i, j, old}); xla::DExpr padded_expr = xla::DExpr::Const(filled_batch); - xla::DExpr subst_expr = e.substitute(1, padded_expr).simplify(); + const std::set ids = e->get_all_ids(); + if (ids.size() != 1) { + return errors::InvalidArgument( + "Dynamic shape padding expected exactly one dynamic " + "variable for argument ", + i, ", dimension ", j, ", but found ", ids.size(), + " variables in expression ", DExprToString(e)); + } + const int substitute_var_id = *ids.begin(); + VLOG(1) << "Calling dynamic expression substitute for compile " + << "argument " << i << " dimension " << j + << " expr=" << DExprToString(e) + << " substitute Var(" << substitute_var_id + << ")=" << filled_batch; + xla::DExpr subst_expr = + e.substitute(substitute_var_id, padded_expr).simplify(); + VLOG(1) << "Dynamic expression substitute for compile argument " + << i << " dimension " << j + << " returned " << DExprToString(subst_expr); if (!subst_expr->is_constant()) { return errors::InvalidArgument( "Dynamic shape padding substitution did not produce an " @@ -958,11 +1562,25 @@ void XlaLocalLaunchBase::ComputeAsync(OpKernelContext* ctx, DoneCallback done) { return; } + bool dynamic_solve_conflict = false; absl::Status status = CompileToLocalExecutable( ctx, function_, /*has_ref_vars=*/has_ref_vars_, platform_info_, xla_compiler_args, DeviceCompileMode::kStrict, - /*may_alias_resource_update=*/true, &client, &compilation_result, + /*may_alias_resource_update=*/true, /*force_static_shapes=*/false, + &dynamic_solve_conflict, &client, &compilation_result, &executable); + if (dynamic_solve_conflict) { + LOG(WARNING) << "Retrying XLA cluster " << function_.name() + << " with concrete static argument shapes"; + std::vector static_args = + BuildStaticCompilationArguments(xla_compiler_args); + status = CompileToLocalExecutable( + ctx, function_, /*has_ref_vars=*/has_ref_vars_, platform_info_, + static_args, DeviceCompileMode::kStrict, + /*may_alias_resource_update=*/true, /*force_static_shapes=*/true, + /*dynamic_solve_conflict=*/nullptr, &client, &compilation_result, + &executable); + } OP_REQUIRES_OK_ASYNC(ctx, status, done); // Continuation of the execution, may be run in a different thread. @@ -1151,6 +1769,27 @@ void XlaCompileOp::Compute(OpKernelContext* ctx) { GetXlaOpsCommonFlags() ->tf_xla_use_device_api.IsEnabledInXlaCompileAndRunForDevice( platform_info_.device_type()); + std::vector compiler_args; + auto compile_static_fallback = [&]() -> absl::Status { + std::vector static_args = + BuildStaticCompilationArguments(compiler_args); + kernel = nullptr; + executable = nullptr; + pjrt_executable = nullptr; + LOG(WARNING) << "Retrying XLA cluster " << function_.name() + << " with concrete static argument shapes"; + if (use_pjrt) { + return CompileToPjRtLoadedExecutable( + *ctx, platform_info_, function_, static_args, compile_mode, + has_ref_vars_, /*may_alias_resource_update=*/false, &kernel, + &pjrt_client, &pjrt_executable); + } + return CompileToLocalExecutable( + ctx, function_, has_ref_vars_, platform_info_, static_args, + compile_mode, /*may_alias_resource_update=*/false, + /*force_static_shapes=*/true, + /*dynamic_solve_conflict=*/nullptr, &client, &kernel, &executable); + }; if (GetXlaOpsCommonFlags()->tf_xla_always_defer_compilation || cannot_compile_cluster) { @@ -1159,13 +1798,14 @@ void XlaCompileOp::Compute(OpKernelContext* ctx) { auto args_and_variables_snapshot = GetXlaCompilerArgsAndSnapshotVariables( resources_, constants_, inputs, ctx); OP_REQUIRES_OK(ctx, args_and_variables_snapshot.status()); - const std::vector& args = - args_and_variables_snapshot->first; + compiler_args = std::move(args_and_variables_snapshot->first); + const std::vector& args = compiler_args; variables_snapshot = std::move(args_and_variables_snapshot->second); // Do not alias resource updates as locking variables in XlaCompile and // unlocking them in XlaRun may lead to deadlocks. absl::Status status; + bool dynamic_solve_conflict = false; if (use_pjrt) { VLOG(2) << "Using PJRT for compilation. Function name: " << function_.name(); @@ -1176,7 +1816,11 @@ void XlaCompileOp::Compute(OpKernelContext* ctx) { } else { status = CompileToLocalExecutable( ctx, function_, has_ref_vars_, platform_info_, args, compile_mode, - /*may_alias_resource_update=*/false, &client, &kernel, &executable); + /*may_alias_resource_update=*/false, /*force_static_shapes=*/false, + &dynamic_solve_conflict, &client, &kernel, &executable); + } + if (dynamic_solve_conflict) { + status = compile_static_fallback(); } if (compile_mode != DeviceCompileMode::kLazy || @@ -1209,6 +1853,35 @@ void XlaCompileOp::Compute(OpKernelContext* ctx) { } } + if ((executable || pjrt_executable) && kernel != nullptr && + GetMarkForCompilationPassFlags()->tf_xla_enable_dynamic_sizes) { + DynamicBatchResolutionResult resolution = + ResolveDynamicBatchSizeFromRuntimeInputs(ctx, *kernel, constants_.size(), + /*log_solves=*/true, + function_.name(), def()); + if (!resolution.can_run) { + const std::string error_message = absl::StrCat( + "Rejecting XLA cluster at compile time because dynamic expressions " + "cannot be solved consistently for this request. ", + resolution.diagnostic); + LOG(WARNING) << error_message; + absl::Status static_status = compile_static_fallback(); + if (!static_status.ok()) { + if (must_compile_) { + OP_REQUIRES_OK(ctx, static_status); + } + LOG(WARNING) << "Static XLA fallback failed for cluster " + << function_.name() << ": " << static_status; + executable = nullptr; + pjrt_executable = nullptr; + kernel = nullptr; + } else { + LOG(INFO) << "Using static XLA fallback for cluster " + << function_.name(); + } + } + } + AllocatorAttributes host_alloc_attrs; host_alloc_attrs.set_gpu_compatible(true); host_alloc_attrs.set_on_host(true); @@ -1268,6 +1941,14 @@ void XlaRunOp::Compute(OpKernelContext* ctx) { const PjRtExecutableClosureStore::KeyT& key = key_tensor.flat()(0); PjRtExecutableClosure closure = PjRtExecutableClosureStore::Global()->Consume(key); + const std::string cluster_name = + closure.compilation_result() != nullptr && + closure.compilation_result()->computation != nullptr + ? closure.compilation_result()->computation->name() + : std::string(""); + VLOG(1) << "Entering XLA cluster: cluster=" << cluster_name + << " op=" << def().name() << " key=" << key + << " mode=pjrt step_id=" << ctx->step_id(); // Fetch inputs from the OpKernelContext. Inputs are the same as the ones // for XlaCompile, except that the must-be-constant inputs that appear in @@ -1304,6 +1985,14 @@ void XlaRunOp::Compute(OpKernelContext* ctx) { XlaExecutableClosure closure = XlaExecutableClosureStore::Global()->Consume(key); + const std::string cluster_name = + closure.compilation_result() != nullptr && + closure.compilation_result()->computation != nullptr + ? closure.compilation_result()->computation->name() + : std::string(""); + VLOG(1) << "Entering XLA cluster: cluster=" << cluster_name + << " op=" << def().name() << " key=" << key + << " mode=local step_id=" << ctx->step_id(); std::shared_ptr allocator = GetAllocator(ctx->device(), GetStream(ctx), platform_info_); XlaComputationLaunchContext launch_context = @@ -1342,71 +2031,31 @@ void XlaRunOp::Compute(OpKernelContext* ctx) { MarkForCompilationPassFlags* flags = GetMarkForCompilationPassFlags(); if (flags->tf_xla_enable_dynamic_sizes) { + DynamicBatchResolutionResult batch_resolution = + ResolveDynamicBatchSizeFromRuntimeInputs( + ctx, *closure.compilation_result(), closure.num_constant_args(), + /*log_solves=*/true, cluster_name, def()); bool is_set = false; - std::set dyn_vals; - const auto* comp_result = closure.compilation_result(); - const int num_constant_args = closure.num_constant_args(); - for (int i = 0; i < comp_result->xla_input_shapes.size(); i++) { - const auto& xla_shape = closure.compilation_result()->xla_input_shapes[i]; - if (!xla_shape.IsArray() || xla_shape.expressions().empty()) continue; - - for (int dim = 0; dim < xla_shape.expressions().size(); dim++) { - const auto& expr = xla_shape.expressions(dim); - if (expr && expr->is_dynamic()) { - xla::DExpr simplified_expr = expr.simplify(); - int input_idx = comp_result->input_mapping[i] - num_constant_args; - if (input_idx < 0 || input_idx >= ctx->num_inputs()) { - VLOG(1) << "Warning: Input index is out of range"; - continue; - } - VLOG(1) << "input shape is " << ctx->input(input_idx).shape() - << ", corresponding xla input shape is " << xla_shape; - int64_t size = ctx->input(input_idx).shape().dim_size(dim); - std::optional dyn_val = - simplified_expr->solve( - size); // TODO: check if the result is correct later. - if (dyn_val.has_value()) { - VLOG(1) << "Found dynamic input. Real size is: " << size - << ", solved dynamic value is " << *dyn_val; - } else { - xla::StringPrinter printer; - simplified_expr->print(&printer); - VLOG(1) << "Warning: Failed to solve the expression " - << std::move(printer).ToString(); - continue; - } - dyn_vals.insert(*dyn_val); - } - } + if (!batch_resolution.can_run) { + LOG(ERROR) << batch_resolution.diagnostic; + ctx->CtxFailure(errors::InvalidArgument(batch_resolution.diagnostic)); + return; } - - if (dyn_vals.size() == 1) { - run_options.set_batch_size(*(dyn_vals.begin())); + if (batch_resolution.has_batch_size) { + VLOG(1) << "Setting run_options.batch_size " + << batch_resolution.diagnostic; + run_options.set_batch_size(batch_resolution.batch_size); is_set = true; - } else { - // Found multiple variables - VLOG(1) << "Warning: Found multiple variables"; } - if (!is_set) { - // TODO: Fallback to BatchSizeResource for now. Remove it later. - BatchSizeResource* bsr = nullptr; - ScopedStepContainer* step_container = ctx->step_container(); - - absl::Status st = step_container->Lookup( - ctx->resource_manager(), BatchSizeResourceName, &bsr); - - if (st.ok()) { - run_options.set_batch_size(bsr->GetBatchSize()); - VLOG(1) << "run_options.batch_size is set to: " - << run_options.batch_size() << ". step_id: " << ctx->step_id(); - bsr->Unref(); - - } else if (IsNotFound(st)) { - VLOG(1) << "Warning: Not found BatchSizeResource in step_container."; - } else { - OP_REQUIRES_OK(ctx, st); - } + LOG(WARNING) << "Entering XLA cluster without run_options.batch_size " + << "being set because " << batch_resolution.diagnostic + << ". op=" << def().name() << " closure_key=" << key + << " step_id=" << ctx->step_id() + << " current_run_options_batch_size=" + << run_options.batch_size() + << " batch_size_resource_not_found=" + << batch_resolution.batch_size_resource_not_found; } } diff --git a/tensorflow/compiler/jit/kernels/xla_ops_internal.h b/tensorflow/compiler/jit/kernels/xla_ops_internal.h new file mode 100644 index 00000000000000..cb710283533418 --- /dev/null +++ b/tensorflow/compiler/jit/kernels/xla_ops_internal.h @@ -0,0 +1,66 @@ +/* Copyright 2026 The TensorFlow Authors. All Rights Reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#ifndef TENSORFLOW_COMPILER_JIT_KERNELS_XLA_OPS_INTERNAL_H_ +#define TENSORFLOW_COMPILER_JIT_KERNELS_XLA_OPS_INTERNAL_H_ + +#include +#include +#include + +#include "absl/types/span.h" +#include "tensorflow/compiler/tf2xla/xla_argument.h" + +namespace tensorflow { +namespace xla_ops_internal { + +struct IgnoredDynamicArgumentOccurrence { + enum class Source { + kShapeDimension, + kConstantValueElement, + }; + + Source source; + int arg_index; + int dim_or_index; + int64_t observed_value; + int64_t solved_value; + std::string expr; +}; + +struct DynamicSolveFilterDecision { + bool can_run = true; + std::string diagnostic; + std::vector ignored_occurrences; +}; + +int GetRuntimeInputIndex(absl::Span input_mapping, + int xla_input_index, int num_constant_args, + bool constants_omitted); + +DynamicSolveFilterDecision AnalyzeIgnoredDynamicArgumentOccurrences( + absl::Span args); + +std::vector BuildStaticCompilationArguments( + absl::Span args); + +void StripIgnoredDynamicArgumentOccurrences( + const std::vector& ignored_occurrences, + std::vector* args); + +} // namespace xla_ops_internal +} // namespace tensorflow + +#endif // TENSORFLOW_COMPILER_JIT_KERNELS_XLA_OPS_INTERNAL_H_ diff --git a/tensorflow/compiler/jit/kernels/xla_ops_test.cc b/tensorflow/compiler/jit/kernels/xla_ops_test.cc new file mode 100644 index 00000000000000..2a513ebc42ab56 --- /dev/null +++ b/tensorflow/compiler/jit/kernels/xla_ops_test.cc @@ -0,0 +1,173 @@ +/* Copyright 2026 The TensorFlow Authors. All Rights Reserved. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#include "tensorflow/compiler/jit/kernels/xla_ops_internal.h" + +#include +#include +#include + +#include "xla/shape_expr.h" +#include "tensorflow/core/framework/tensor_shape.h" +#include "tensorflow/core/platform/test.h" + +namespace tensorflow { +namespace xla_ops_internal { +namespace { + +TEST(XlaInputMappingTest, UsesCompilationInputMapping) { + const std::vector input_mapping = {1, 3}; + + EXPECT_EQ(GetRuntimeInputIndex(input_mapping, 0, /*num_constant_args=*/1, + /*constants_omitted=*/false), + 1); + EXPECT_EQ(GetRuntimeInputIndex(input_mapping, 1, /*num_constant_args=*/1, + /*constants_omitted=*/false), + 3); + EXPECT_EQ(GetRuntimeInputIndex(input_mapping, 0, /*num_constant_args=*/1, + /*constants_omitted=*/true), + 0); + EXPECT_EQ(GetRuntimeInputIndex(input_mapping, 1, /*num_constant_args=*/1, + /*constants_omitted=*/true), + 2); + EXPECT_EQ(GetRuntimeInputIndex(input_mapping, 2, /*num_constant_args=*/1, + /*constants_omitted=*/true), + -1); +} + +XlaArgument MakeDynamicArgument(int64_t observed_size, + xla::DExpr expr = xla::DExpr::Var(1)) { + XlaArgument arg; + arg.kind = XlaArgument::kParameter; + arg.type = DT_FLOAT; + TensorShape shape({observed_size, 16}); + shape.set_expression(0, std::move(expr)); + shape.set_expression(1, xla::DExpr::Const(16)); + arg.shape = shape; + return arg; +} + +TEST(DynamicSolveFilterTest, IgnoresSingletonWhenNonSingletonEvidenceExists) { + std::vector args = { + MakeDynamicArgument(1), + MakeDynamicArgument(240), + }; + + DynamicSolveFilterDecision decision = + AnalyzeIgnoredDynamicArgumentOccurrences(args); + + ASSERT_TRUE(decision.can_run); + ASSERT_EQ(decision.ignored_occurrences.size(), 1); + EXPECT_EQ(decision.ignored_occurrences[0].arg_index, 0); + EXPECT_EQ(decision.ignored_occurrences[0].observed_value, 1); + + StripIgnoredDynamicArgumentOccurrences(decision.ignored_occurrences, &args); + const TensorShape& singleton_shape = std::get(args[0].shape); + const TensorShape& nonsingleton_shape = std::get(args[1].shape); + EXPECT_FALSE(singleton_shape.get_expression(0)); + EXPECT_EQ(singleton_shape.dim_size(0), 1); + EXPECT_TRUE(nonsingleton_shape.get_expression(0)); +} + +TEST(DynamicSolveFilterTest, + IgnoresSingletonAcrossExpressionsOfTheSameVariable) { + const xla::DExpr variable = xla::DExpr::Var(1); + std::vector args = { + MakeDynamicArgument(1, variable), + MakeDynamicArgument(241, variable + 1), + }; + + DynamicSolveFilterDecision decision = + AnalyzeIgnoredDynamicArgumentOccurrences(args); + + ASSERT_TRUE(decision.can_run); + ASSERT_EQ(decision.ignored_occurrences.size(), 1); + EXPECT_EQ(decision.ignored_occurrences[0].arg_index, 0); + EXPECT_EQ(decision.ignored_occurrences[0].observed_value, 1); + EXPECT_EQ(decision.ignored_occurrences[0].solved_value, 1); +} + +TEST(DynamicSolveFilterTest, RejectsConflictingNonSingletonEvidence) { + std::vector args = { + MakeDynamicArgument(240), + MakeDynamicArgument(720), + }; + + DynamicSolveFilterDecision decision = + AnalyzeIgnoredDynamicArgumentOccurrences(args); + + EXPECT_FALSE(decision.can_run); + EXPECT_TRUE(decision.ignored_occurrences.empty()); + EXPECT_NE(decision.diagnostic.find("conflicting non-singleton candidates"), + std::string::npos); +} + +TEST(DynamicSolveFilterTest, KeepsConsistentSingletonOnlyEvidence) { + std::vector args = { + MakeDynamicArgument(1), + MakeDynamicArgument(1), + }; + + DynamicSolveFilterDecision decision = + AnalyzeIgnoredDynamicArgumentOccurrences(args); + + EXPECT_TRUE(decision.can_run); + EXPECT_TRUE(decision.ignored_occurrences.empty()); +} + +TEST(DynamicSolveFilterTest, StaticArgumentsKeepConcreteRuntimeShapes) { + XlaArgument dynamic_arg = MakeDynamicArgument(37); + dynamic_arg.constant_value_expressions.resize(1); + dynamic_arg.constant_value_expressions[0].set_variable_id(1); + + std::vector args = {dynamic_arg}; + std::vector static_args = + BuildStaticCompilationArguments(args); + + ASSERT_EQ(static_args.size(), 1); + const TensorShape& static_shape = + std::get(static_args[0].shape); + EXPECT_EQ(static_shape.dim_size(0), 37); + EXPECT_EQ(static_shape.dim_size(1), 16); + EXPECT_TRUE(static_shape.get_expressions().empty()); + EXPECT_TRUE(static_args[0].constant_value_expressions.empty()); +} + +TEST(DynamicSolveFilterTest, StaticArgumentsClearXlaShapeDynamism) { + xla::Shape dynamic_shape; + dynamic_shape.set_element_type(xla::F32); + dynamic_shape.add_dimensions(37, /*is_dynamic=*/true, + xla::DExpr::Var(1)); + + XlaArgument dynamic_arg; + dynamic_arg.kind = XlaArgument::kParameter; + dynamic_arg.type = DT_FLOAT; + dynamic_arg.shape = dynamic_shape; + std::vector args = {dynamic_arg}; + + std::vector static_args = + BuildStaticCompilationArguments(args); + + const xla::Shape& static_shape = std::get(static_args[0].shape); + EXPECT_EQ(static_shape.dimensions(0), 37); + EXPECT_FALSE(static_shape.is_dynamic_dimension(0)); + ASSERT_EQ(static_shape.expressions().size(), 1); + EXPECT_TRUE(static_shape.expressions(0)->is_constant()); + EXPECT_EQ(static_shape.expressions(0)->get_val(), 37); +} + +} // namespace +} // namespace xla_ops_internal +} // namespace tensorflow diff --git a/tensorflow/compiler/jit/mark_for_compilation_pass.cc b/tensorflow/compiler/jit/mark_for_compilation_pass.cc index 566bab23a11867..81365b13e5a9d5 100644 --- a/tensorflow/compiler/jit/mark_for_compilation_pass.cc +++ b/tensorflow/compiler/jit/mark_for_compilation_pass.cc @@ -267,10 +267,18 @@ class MarkForCompilationPassImpl { dim_vars_.insert(dim_vars.begin(), dim_vars.end()); } const std::set& dim_vars() const { return dim_vars_; } + void add_dim_expr(const xla::DExpr& dim_expr) { + dim_exprs_.push_back(dim_expr); + } + void merge_dim_exprs(const std::vector& dim_exprs) { + dim_exprs_.insert(dim_exprs_.end(), dim_exprs.begin(), dim_exprs.end()); + } + const std::vector& dim_exprs() const { return dim_exprs_; } private: int annotated_id_ = -1; std::set dim_vars_; + std::vector dim_exprs_; int chain_id_ = -1; int cluster_size_ = 1; int cycles_graph_node_id_; @@ -345,6 +353,10 @@ class MarkForCompilationPassImpl { absl::Status AssignAnnotatedClusterIDs(); absl::Status AssignDimVars(); + std::optional CheckDynamicExpressionCompatibility( + absl::Span exprs); + bool DynamicNodeExpressionsAreCompatible(const Node& node, + std::string* reason); void collectInputNodes(std::set &path_nodes); void collectMergeNodes(const std::vector& nodeSet, std::set &merger_nodes); @@ -645,6 +657,8 @@ void MarkForCompilationPassImpl::Cluster::Merge(Cluster* other) { merge_dim_vars(other->dim_vars_); other->dim_vars_.clear(); + merge_dim_exprs(other->dim_exprs_); + other->dim_exprs_.clear(); resource_var_operation_node_ids_.reserve( resource_var_operation_node_ids_.size() + @@ -714,73 +728,186 @@ 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 ""; } } -std::unique_ptr ExprFromProto(const ExpressionProto& proto) { - switch (proto.node_type_case()) { - case ExpressionProto::kConstantValue: - return DimExpr::Cons(proto.constant_value()); - case ExpressionProto::kVariableId: - return DimExpr::Var(proto.variable_id()); - case ExpressionProto::kAddNode: { - auto lhs = ExprFromProto(proto.add_node().lhs()); - auto rhs = ExprFromProto(proto.add_node().rhs()); - // Note: These are owning pointers, but ExprAdd takes raw pointers. - // The caller must manage lifetime appropriately. - return std::make_unique(lhs.release(), rhs.release()); - } - case ExpressionProto::kSubNode: { - auto lhs = ExprFromProto(proto.sub_node().lhs()); - auto rhs = ExprFromProto(proto.sub_node().rhs()); - return std::make_unique(lhs.release(), rhs.release()); - } - case ExpressionProto::kMulNode: { - auto lhs = ExprFromProto(proto.mul_node().lhs()); - auto rhs = ExprFromProto(proto.mul_node().rhs()); - return std::make_unique(lhs.release(), rhs.release()); - } - case ExpressionProto::kDivNode: { - auto lhs = ExprFromProto(proto.div_node().lhs()); - auto rhs = ExprFromProto(proto.div_node().rhs()); - return std::make_unique(lhs.release(), rhs.release()); - } - case ExpressionProto::NODE_TYPE_NOT_SET: - default: - return nullptr; +bool HasDynamicInputExpression(const Node* node) { + for (const Edge* edge : node->in_edges()) { + if (edge->IsControlEdge()) { + continue; + } + auto it = expr_map.find(edge->src()->name()); + if (it == expr_map.end()) { + continue; + } + const int output_index = edge->src_output(); + if (output_index < 0 || output_index >= it->second.size()) { + continue; + } + for (const auto& expr : it->second[output_index]) { + if (expr == nullptr) { + continue; + } + xla::DExpr dynamic_expr = *expr; + if (dynamic_expr && dynamic_expr->is_dynamic()) { + return true; + } + } } + return false; } -static xla::DExpr DimExprToDExpr(const DimExpr* e) { - switch (e->kind()) { - case DimExpr::Kind::kConstant: { - auto* ac = static_cast(e); - return xla::DExpr::Const(ac->value()); +std::string DynamicExpressionToString(const xla::DExpr& expr) { + xla::DExpr simplified = expr.simplify(); + if (!simplified && !simplified.is_unknown()) { + return ""; + } + xla::StringPrinter printer; + simplified->print(&printer); + return std::move(printer).ToString(); +} + +std::optional CheckDynamicExpressionCompatibilityImpl( + absl::Span exprs) { + const xla::DExpr* anchor_source = nullptr; + for (const xla::DExpr& expr : exprs) { + if (expr && expr->is_dynamic() && !expr->get_all_ids().empty()) { + anchor_source = &expr; + break; } - case DimExpr::Kind::kVariable: { - auto* av = static_cast(e); - return xla::DExpr::Var(av->id()); // Use 1 all the time for now + } + if (anchor_source == nullptr) { + return std::nullopt; + } + + const std::set expected_ids = (*anchor_source)->get_all_ids(); + std::optional expected_core; + if (expected_ids.size() > 1) { + expected_core = + anchor_source->find_smallest_subexpression_covering_all_variables(); + } + int replacement_id = std::numeric_limits::min(); + while (expected_ids.count(replacement_id) != 0) { + ++replacement_id; + } + const std::set replacement_ids = {replacement_id}; + + for (const xla::DExpr& expr : exprs) { + if (!expr || !expr->is_dynamic() || expr->get_all_ids().empty()) { + continue; } - case DimExpr::Kind::kAdd: { - auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) + DimExprToDExpr(ee->rhs()); + const std::set expr_ids = expr->get_all_ids(); + if (expr_ids != expected_ids) { + return absl::StrCat( + "dynamic expressions do not share the same variable ids: " + "expected={", + absl::StrJoin(expected_ids, ", "), "}, expr=", + DynamicExpressionToString(expr), ", ids={", + absl::StrJoin(expr_ids, ", "), "}"); } - case DimExpr::Kind::kSub: { - auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) - DimExprToDExpr(ee->rhs()); + if (!expected_core.has_value()) { + continue; } - case DimExpr::Kind::kMul: { - auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) * DimExprToDExpr(ee->rhs()); + const xla::DExpr core = + expr.find_smallest_subexpression_covering_all_variables(); + if (!(core == *expected_core)) { + return absl::StrCat( + "dynamic expressions do not share the same clusterable core: " + "expected=", + DynamicExpressionToString(*expected_core), ", expr=", + DynamicExpressionToString(expr), ", core=", + DynamicExpressionToString(core)); } - case DimExpr::Kind::kDiv: { - auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) / DimExprToDExpr(ee->rhs()); + const xla::DExpr normalized = + expr.replace_subexpression(core, xla::DExpr::Var(replacement_id)) + .simplify(); + if (normalized->get_all_ids() != replacement_ids) { + return absl::StrCat( + "dynamic expression core does not cover every variable occurrence: " + "expr=", + DynamicExpressionToString(expr), ", core=", + DynamicExpressionToString(core)); } } - return xla::DExpr(); + + return std::nullopt; +} + +std::optional +MarkForCompilationPassImpl::CheckDynamicExpressionCompatibility( + absl::Span exprs) { + return CheckDynamicExpressionCompatibilityImpl(exprs); +} + +bool MarkForCompilationPassImpl::DynamicNodeExpressionsAreCompatible( + const Node& node, std::string* reason) { + std::vector node_exprs; + + for (const Edge* edge : node.in_edges()) { + if (edge->IsControlEdge()) { + continue; + } + + const Node* src = edge->src(); + auto it = expr_map.find(src->name()); + if (it == expr_map.end()) { + continue; + } + + const int output_index = edge->src_output(); + if (output_index < 0 || + output_index >= static_cast(it->second.size())) { + continue; + } + + for (const auto& expr_ptr : it->second[output_index]) { + if (expr_ptr == nullptr) { + continue; + } + xla::DExpr dyn = *expr_ptr; + if (!dyn || !dyn->is_dynamic() || dyn->get_all_ids().empty()) { + continue; + } + node_exprs.push_back(std::move(dyn)); + } + } + + auto output_it = expr_map.find(node.name()); + if (output_it != expr_map.end()) { + for (const auto& output_exprs : output_it->second) { + for (const auto& expr_ptr : output_exprs) { + if (expr_ptr == nullptr) { + continue; + } + xla::DExpr dyn = *expr_ptr; + if (!dyn || !dyn->is_dynamic() || dyn->get_all_ids().empty()) { + continue; + } + node_exprs.push_back(std::move(dyn)); + } + } + } + + const std::optional incompatibility = + CheckDynamicExpressionCompatibility(node_exprs); + if (!incompatibility.has_value()) { + return true; + } + *reason = *incompatibility; + return false; } // Runs Grappler static inference and logs any ExpressionProto found in output @@ -793,6 +920,8 @@ void LogExpressionsViaGraphProperties(tensorflow::Graph& graph) { using tensorflow::grappler::GraphProperties; using tensorflow::grappler::GrapplerItem; + expr_map.clear(); + GraphDef graph_def; graph.ToGraphDef(&graph_def); auto node_name_index = graph.BuildNodeNameIndex(); @@ -807,7 +936,8 @@ void LogExpressionsViaGraphProperties(tensorflow::Graph& graph) { /*assume_valid_feeds=*/false, /*aggressive_shape_inference=*/false, /*include_input_tensor_values=*/false, - /*include_output_tensor_values=*/false); + /*include_output_tensor_values=*/false, + /*enable_dynamic_value_inference=*/true); if (!st.ok()) { LOG(ERROR) << "[EXPR][GP] InferStatically failed: " << st.message(); @@ -820,11 +950,14 @@ void LogExpressionsViaGraphProperties(tensorflow::Graph& graph) { auto convert_graph_properties_shape = [](const TensorShapeProto& gp_shape) { TensorShapeProto out; out.set_unknown_rank(gp_shape.unknown_rank()); - for (const auto& dim : gp_shape.dim()) { + for (int i = 0; i < gp_shape.dim_size(); ++i) { + const auto& dim = gp_shape.dim(i); out.add_dim()->set_size(dim.size()); ExpressionProto* expr = out.add_expressions(); - if (dim.expr().node_type_case() != ExpressionProto::NODE_TYPE_NOT_SET) { - *expr = dim.expr(); + if (i < gp_shape.expressions_size() && + gp_shape.expressions(i).node_type_case() != + ExpressionProto::NODE_TYPE_NOT_SET) { + *expr = gp_shape.expressions(i); } else { expr->set_constant_value(dim.size()); } @@ -845,24 +978,23 @@ void LogExpressionsViaGraphProperties(tensorflow::Graph& graph) { std::vector> exprs; for (int d = 0; d < shp.dim_size(); ++d) { - const auto& dim = shp.dim(d); - - const ExpressionProto& expr = dim.expr(); + if (d >= shp.expressions_size()) continue; + const ExpressionProto& expr = shp.expressions(d); if (expr.node_type_case() == ExpressionProto::NODE_TYPE_NOT_SET) continue; VLOG(1) << "Node " << n.name() << " has expression " << ExprProtoToString(expr); - auto ex = ExprFromProto(expr); - exprs.push_back(std::move(ex)); + exprs.push_back( + std::make_unique(DimExprFromProto(expr))); ++found; } if (shp.dim_size() == 0 && shp.unknown_rank()) { // Add two dummy variables to represent the unknown rank - exprs.push_back(std::make_unique(-888)); - exprs.push_back(std::make_unique(-889)); + exprs.push_back(std::make_unique(DimExpr::Var(-888))); + exprs.push_back(std::make_unique(DimExpr::Var(-889))); } list_exprs[out_idx] = std::move(exprs); @@ -884,6 +1016,9 @@ absl::StatusOr MarkForCompilationPassImpl::Initialize() { TF_RET_CHECK(!initialized_ && !edges_contracted_ && !clusters_created_); initialized_ = true; + if (debug_options_.enable_dynamic_sizes) { + LogExpressionsViaGraphProperties(*graph_); + } TF_RETURN_IF_ERROR(FindCompilationCandidates()); if (compilation_candidates_.empty()) { @@ -921,47 +1056,22 @@ absl::StatusOr MarkForCompilationPassImpl::Initialize() { TF_RETURN_IF_ERROR(AssignAnnotatedClusterIDs()); } if (debug_options_.enable_dynamic_sizes) { - LogExpressionsViaGraphProperties(*graph_); TF_RETURN_IF_ERROR(AssignDimVars()); - auto has_dynamic_input_expression = [&](const Node* n) { - for (const Edge* edge : n->in_edges()) { - if (edge->IsControlEdge()) { - continue; - } - const Node* src = edge->src(); - auto it = expr_map.find(src->name()); - if (it == expr_map.end()) { - continue; - } - const int output_index = edge->src_output(); - if (output_index < 0 || output_index >= it->second.size()) { - continue; - } - for (const auto& expr_ptr : it->second[output_index]) { - if (expr_ptr == nullptr) { - continue; - } - xla::DExpr dyn = DimExprToDExpr(expr_ptr.get()); - if (dyn && dyn->is_dynamic()) { - return true; - } - } - } - return false; - }; for (Node* n : graph_->op_nodes()) { bool mark_shape_derived = false; - if (n->type_string() == "Shape" || n->type_string() == "ShapeN") { - mark_shape_derived = has_dynamic_input_expression(n); + auto is_shape_like = [](const Node* node) { + const string& type = node->type_string(); + return type == "Shape" || type == "ShapeN" || type == "Size"; + }; + if (is_shape_like(n)) { + mark_shape_derived = HasDynamicInputExpression(n); } else if (n->type_string() == "Cast") { for (const Edge* edge : n->in_edges()) { if (edge->IsControlEdge()) { continue; } const Node* src = edge->src(); - if ((src->type_string() == "Shape" || - src->type_string() == "ShapeN") && - has_dynamic_input_expression(src)) { + if (is_shape_like(src) && HasDynamicInputExpression(src)) { mark_shape_derived = true; break; } @@ -1726,6 +1836,15 @@ absl::Status MarkForCompilationPassImpl::FindCompilationCandidates() { continue; } + if (debug_options_.enable_dynamic_sizes) { + std::string reason; + if (!DynamicNodeExpressionsAreCompatible(*node, &reason)) { + VLOG(1) << "Rejecting " << node->name() + << " from XLA clustering: " << reason; + continue; + } + } + if (compile_time_const_nodes[node->id()]) { const OpDef* op_def; TF_RETURN_IF_ERROR( @@ -1909,7 +2028,11 @@ absl::Status MarkForCompilationPassImpl::AssignDimVars(void) { } for (auto& pDim: (it->second)[output_index]) { DimExpr * d= pDim.get(); - xla::DExpr dyn = DimExprToDExpr(d); + xla::DExpr dyn = *d; + if (!dyn || !dyn->is_dynamic()) { + continue; + } + cluster->add_dim_expr(dyn); auto new_ids = dyn->get_all_ids(); for (auto id : new_ids) { cluster->add_dim_var(id); @@ -1917,6 +2040,21 @@ absl::Status MarkForCompilationPassImpl::AssignDimVars(void) { } } } + auto output_it = expr_map.find(node_name); + if (output_it != expr_map.end()) { + for (const auto& output_exprs : output_it->second) { + for (const auto& expr_ptr : output_exprs) { + if (expr_ptr == nullptr) { + continue; + } + xla::DExpr dyn = *expr_ptr; + if (!dyn || !dyn->is_dynamic()) { + continue; + } + cluster->add_dim_expr(dyn); + } + } + } // create a for loop for each dim vars in cluster and print each dim var if (VLOG_IS_ON(2)) { if (cluster->dim_vars().empty()) { @@ -2161,29 +2299,13 @@ absl::StatusOr MarkForCompilationPassImpl::TryToContractEdge( } if (debug_options_.enable_dynamic_sizes) { - if (from->dim_vars().size() > 1 || to->dim_vars().size() > 1) { - std::string from_str = "from_vars: "; - for (auto id : from->dim_vars()) { - from_str += std::to_string(id) + ", "; - } - std::string to_str = "to_vars: "; - for (auto id : to->dim_vars()) { - to_str += std::to_string(id) + ", "; - } - return LogNotContractableAndReturnFalse( - from, to, absl::StrCat("the two nodes have multiple dynamic dimensions: ", - from_str, " and ", to_str)); - } - if (from->dim_vars().size() == 1 && to->dim_vars().size() == 1 && - from->dim_vars() != to->dim_vars()) { - return LogNotContractableAndReturnFalse( - from, to, - absl::StrCat("the two nodes have different dynamic dimensions: ", - from->dim_vars().size() == 1 - ? std::to_string(*from->dim_vars().begin()) : "none", - " and ", - to->dim_vars().size() == 1 - ? std::to_string(*to->dim_vars().begin()) : "none")); + std::vector combined_exprs = from->dim_exprs(); + combined_exprs.insert(combined_exprs.end(), to->dim_exprs().begin(), + to->dim_exprs().end()); + std::optional incompatibility = + CheckDynamicExpressionCompatibility(combined_exprs); + if (incompatibility.has_value()) { + return LogNotContractableAndReturnFalse(from, to, *incompatibility); } } @@ -2679,6 +2801,11 @@ void ResetClusterSequenceNumber() { ClusterSequenceNumberGenerator::Global().Reset(); } +std::optional CheckDynamicExpressionCompatibilityForTest( + absl::Span exprs) { + return CheckDynamicExpressionCompatibilityImpl(exprs); +} + absl::flat_hash_set GetKnownXLAAllowlistOp() { absl::flat_hash_set result{ "AdjustContrastv2", diff --git a/tensorflow/compiler/jit/mark_for_compilation_pass.h b/tensorflow/compiler/jit/mark_for_compilation_pass.h index 558912f2eee2e0..79e3c88074fcdf 100644 --- a/tensorflow/compiler/jit/mark_for_compilation_pass.h +++ b/tensorflow/compiler/jit/mark_for_compilation_pass.h @@ -20,10 +20,17 @@ limitations under the License. #ifndef TENSORFLOW_COMPILER_JIT_MARK_FOR_COMPILATION_PASS_H_ #define TENSORFLOW_COMPILER_JIT_MARK_FOR_COMPILATION_PASS_H_ +#include + #include "absl/container/flat_hash_set.h" +#include "absl/types/span.h" #include "tensorflow/compiler/jit/compilability_check_util.h" #include "tensorflow/core/common_runtime/optimization_registry.h" +namespace xla { +class DExpr; +} + namespace tensorflow { // The attribute that marks nodes to be grouped into functions by the @@ -57,6 +64,11 @@ void ResetClusterSequenceNumber(); // Return a list of operation that we choose not to put into the allowlist. absl::flat_hash_set GetKnownXLAAllowlistOp(); + +// Returns an explanation when dynamic expressions cannot use the same +// cluster-level dynamic value. +std::optional CheckDynamicExpressionCompatibilityForTest( + absl::Span exprs); } // namespace testing } // namespace tensorflow diff --git a/tensorflow/compiler/jit/mark_for_compilation_pass_test.cc b/tensorflow/compiler/jit/mark_for_compilation_pass_test.cc index 1a120791206369..bdc4344af22547 100644 --- a/tensorflow/compiler/jit/mark_for_compilation_pass_test.cc +++ b/tensorflow/compiler/jit/mark_for_compilation_pass_test.cc @@ -53,6 +53,7 @@ limitations under the License. #include "tensorflow/core/lib/core/status_test_util.h" #include "tensorflow/core/platform/errors.h" #include "tensorflow/core/platform/test.h" +#include "xla/shape_expr.h" using ::tensorflow::testing::FindNodeByName; @@ -117,6 +118,74 @@ absl::flat_hash_map> GetClusterSets( return cluster_sets; } +// Expressions may differ outside the smallest subtree containing all variables. +TEST(XlaCompilationTest, + DynamicExpressionCompatibilityUsesSmallestCoveringSubexpression) { + const xla::DExpr shared_core = + xla::DExpr::Var(1) + xla::DExpr::Var(2); + const std::vector compatible_exprs = { + shared_core + 2, shared_core + 3}; + + EXPECT_FALSE( + testing::CheckDynamicExpressionCompatibilityForTest(compatible_exprs) + .has_value()); + + const std::vector incompatible_exprs = { + shared_core + 2, + (xla::DExpr::Var(1) - xla::DExpr::Var(2)) + 3}; + const std::optional incompatibility = + testing::CheckDynamicExpressionCompatibilityForTest(incompatible_exprs); + + ASSERT_TRUE(incompatibility.has_value()); + EXPECT_NE(incompatibility->find("same clusterable core"), + std::string::npos); +} + +TEST(XlaCompilationTest, + DynamicExpressionCompatibilityRejectsDifferentVariableSets) { + const std::vector exprs = { + xla::DExpr::Var(1) + xla::DExpr::Var(2), + xla::DExpr::Var(1) + xla::DExpr::Var(3)}; + + const std::optional incompatibility = + testing::CheckDynamicExpressionCompatibilityForTest(exprs); + + ASSERT_TRUE(incompatibility.has_value()); + EXPECT_NE(incompatibility->find("same variable ids"), std::string::npos); +} + +TEST(XlaCompilationTest, + DynamicExpressionCompatibilityUsesCompleteRepeatedVariableCore) { + const xla::DExpr a = xla::DExpr::Var(1); + const xla::DExpr b = xla::DExpr::Var(2); + const xla::DExpr shared_core = (a + b) * (a - b); + const std::vector exprs = {shared_core + 2, shared_core + 3}; + + EXPECT_FALSE( + testing::CheckDynamicExpressionCompatibilityForTest(exprs).has_value()); +} + +TEST(XlaCompilationTest, + DynamicExpressionCompatibilityAcceptsSharedCoreWithDifferentScaling) { + const xla::DExpr shared_core = + xla::DExpr::Var(1) + xla::DExpr::Var(2); + const std::vector exprs = { + shared_core / 240, shared_core, xla::DExpr::Const(3) * shared_core}; + + EXPECT_FALSE( + testing::CheckDynamicExpressionCompatibilityForTest(exprs).has_value()); +} + +TEST(XlaCompilationTest, + DynamicExpressionCompatibilityAcceptsSingleVariableExpressions) { + const std::vector exprs = { + xla::DExpr::Var(1) + 2, + xla::DExpr::Var(1) * xla::DExpr::Const(3)}; + + EXPECT_FALSE( + testing::CheckDynamicExpressionCompatibilityForTest(exprs).has_value()); +} + TEST(XlaCompilationTest, Chains) { std::unique_ptr graph(new Graph(OpRegistry::Global())); { diff --git a/tensorflow/compiler/jit/xla_launch_util.cc b/tensorflow/compiler/jit/xla_launch_util.cc index ea2a50e241ebe2..b6b1f08fe73dcb 100644 --- a/tensorflow/compiler/jit/xla_launch_util.cc +++ b/tensorflow/compiler/jit/xla_launch_util.cc @@ -460,8 +460,36 @@ absl::Status XlaComputationLaunchContext::PopulateOutputs( has_dynamic = true; VLOG(1) << "Current expression is " << expr; if (run_options) { - xla::DExpr batch_size = xla::DExpr::Const(run_options->batch_size()); - xla::DExpr subst_expr = expr.substitute(1, batch_size).simplify(); + const int64_t run_options_batch_size = run_options->batch_size(); + if (run_options_batch_size <= 0) { + return absl::InvalidArgumentError(absl::StrCat( + "Cannot substitute dynamic output shape for output ", i, + ", dimension ", dim, + ": the XLA runtime batch size was not initialized")); + } + VLOG(1) << "PopulateOutputs read run_options->batch_size()=" + << run_options_batch_size << " for output " << i + << " dimension " << dim; + xla::DExpr batch_size = xla::DExpr::Const(run_options_batch_size); + const std::set ids = expr->get_all_ids(); + if (ids.size() != 1) { + return absl::InvalidArgumentError(absl::StrCat( + "Runtime shape substitution expected exactly one dynamic " + "variable for output ", + i, ", dimension ", dim, ", but found ", ids.size(), + " variables in expression ", DExprToString(expr))); + } + const int substitute_var_id = *ids.begin(); + VLOG(1) << "Calling output shape substitute with run_options for " + << "output " << i << " dimension " << dim + << " expr=" << DExprToString(expr) + << " substitute Var(" << substitute_var_id << ")=" + << DExprToString(batch_size); + xla::DExpr subst_expr = + expr.substitute(substitute_var_id, batch_size).simplify(); + VLOG(1) << "Output shape substitute with run_options for output " + << i << " dimension " << dim << " returned " + << DExprToString(subst_expr); if (!subst_expr->is_constant()) { return absl::InvalidArgumentError(absl::StrCat( "Runtime shape substitution did not produce an integer " @@ -471,14 +499,46 @@ absl::Status XlaComputationLaunchContext::PopulateOutputs( shape.set_dim(dim, subst_expr->get_val()); } else { // TODO: Fallback to BatchSizeResource for now. Remove it later. - VLOG(1) << "Warning: Didn't find run_options"; + LOG(WARNING) << "PopulateOutputs did not receive run_options for " + << "output " << i << " dimension " << dim + << "; falling back to BatchSizeResource"; BatchSizeResource* bsr = nullptr; ScopedStepContainer* step_container = ctx->step_container(); TF_RETURN_IF_ERROR(step_container->Lookup( ctx->resource_manager(), BatchSizeResourceName, &bsr)); - xla::DExpr batch_size = xla::DExpr::Const(bsr->GetBatchSize()); - // Just substitute Var(1) for now. - xla::DExpr subst_expr = expr.substitute(1, batch_size).simplify(); + if (bsr == nullptr) { + return errors::Internal( + "BatchSizeResource lookup succeeded but returned null"); + } + core::ScopedUnref bsr_ref(bsr); + const int64_t runtime_batch_size = bsr->GetBatchSize(); + if (runtime_batch_size <= 0) { + return absl::InvalidArgumentError(absl::StrCat( + "Cannot substitute dynamic output shape for output ", i, + ", dimension ", dim, + ": the XLA runtime batch size was not initialized")); + } + xla::DExpr batch_size = xla::DExpr::Const(runtime_batch_size); + const std::set ids = expr->get_all_ids(); + if (ids.size() != 1) { + return absl::InvalidArgumentError(absl::StrCat( + "Runtime shape substitution expected exactly one dynamic " + "variable for output ", + i, ", dimension ", dim, ", but found ", ids.size(), + " variables in expression ", DExprToString(expr))); + } + const int substitute_var_id = *ids.begin(); + VLOG(1) << "Calling output shape substitute with " + << "BatchSizeResource for output " << i + << " dimension " << dim + << " expr=" << DExprToString(expr) + << " substitute Var(" << substitute_var_id + << ")=" << DExprToString(batch_size); + xla::DExpr subst_expr = + expr.substitute(substitute_var_id, batch_size).simplify(); + VLOG(1) << "Output shape substitute with BatchSizeResource for " + << "output " << i << " dimension " << dim + << " returned " << DExprToString(subst_expr); if (!subst_expr->is_constant()) { return absl::InvalidArgumentError(absl::StrCat( "Runtime shape substitution did not produce an integer " @@ -486,7 +546,6 @@ absl::Status XlaComputationLaunchContext::PopulateOutputs( i, ", dimension ", dim, ": ", DExprToString(subst_expr))); } shape.set_dim(dim, subst_expr->get_val()); - bsr->Unref(); } } } diff --git a/tensorflow/compiler/tf2xla/BUILD b/tensorflow/compiler/tf2xla/BUILD index c571f73beff0be..52deb99d5c771b 100644 --- a/tensorflow/compiler/tf2xla/BUILD +++ b/tensorflow/compiler/tf2xla/BUILD @@ -1565,8 +1565,11 @@ cc_library( srcs = ["mlir_xla_op_kernel.cc"], hdrs = ["mlir_xla_op_kernel.h"], deps = [ + ":common", ":xla_compiler", ":xla_expression", + "//tensorflow/compiler/jit:flags", + "//tensorflow/compiler/jit:shape_inference", "//tensorflow/compiler/jit:xla_compile_util", "//tensorflow/compiler/mlir/tf2xla/api/v1:compile_mlir_util_no_tf_dialect_passes", "//tensorflow/compiler/mlir/utils:array_container_utils", diff --git a/tensorflow/compiler/tf2xla/kernels/batch_norm_op.cc b/tensorflow/compiler/tf2xla/kernels/batch_norm_op.cc index 0dd528e3dea173..1db3ab4c78c461 100644 --- a/tensorflow/compiler/tf2xla/kernels/batch_norm_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/batch_norm_op.cc @@ -21,9 +21,9 @@ limitations under the License. #include #include "tensorflow/compiler/tf2xla/kernels/relu_op.h" -#include "tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h" #include "tensorflow/compiler/tf2xla/type_util.h" #include "tensorflow/compiler/tf2xla/xla_helpers.h" +#include "tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h" #include "tensorflow/compiler/tf2xla/xla_op_kernel.h" #include "tensorflow/compiler/tf2xla/xla_op_registry.h" #include "xla/hlo/builder/lib/constants.h" @@ -241,7 +241,9 @@ class FusedBatchNormOpEx : public FusedBatchNormOp { REGISTER_XLA_OP(Name("FusedBatchNorm"), FusedBatchNormOp); REGISTER_XLA_OP(Name("FusedBatchNormV2"), FusedBatchNormOp); -REGISTER_XLA_OP(Name("FusedBatchNormV3"), MlirXlaOpKernel); +REGISTER_XLA_OP_FACTORY( + Name("FusedBatchNormV3"), + CreateDynamicNativeXlaOpKernel); REGISTER_XLA_OP(Name("_FusedBatchNormEx"), FusedBatchNormOpEx); class FusedBatchNormGradOp : public XlaOpKernel { @@ -358,7 +360,9 @@ class FusedBatchNormGradOp : public XlaOpKernel { REGISTER_XLA_OP(Name("FusedBatchNormGrad"), FusedBatchNormGradOp); REGISTER_XLA_OP(Name("FusedBatchNormGradV2"), FusedBatchNormGradOp); -REGISTER_XLA_OP(Name("FusedBatchNormGradV3"), MlirXlaOpKernel); +REGISTER_XLA_OP_FACTORY( + Name("FusedBatchNormGradV3"), + CreateDynamicNativeXlaOpKernel); } // namespace } // namespace tensorflow 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/binary_ops.cc b/tensorflow/compiler/tf2xla/kernels/binary_ops.cc index ab9b09a7ea3f30..34762dfeb9b545 100644 --- a/tensorflow/compiler/tf2xla/kernels/binary_ops.cc +++ b/tensorflow/compiler/tf2xla/kernels/binary_ops.cc @@ -48,14 +48,17 @@ namespace { explicit NAME##Op(OpKernelConstruction* ctx) : XlaBinaryOp(ctx) {} \ xla::XlaOp Computation( \ XlaOpKernelContext* ctx, const xla::XlaOp& lhs, \ - const absl::Span& lhs_shape, const xla::XlaOp& rhs, \ + const absl::Span& lhs_shape, \ + const xla::XlaOp& rhs, \ const absl::Span& rhs_shape, \ const BCast& broadcast_helper, \ + const absl::Span& broadcast_output_exprs, \ const std::vector& extend_dimensions) override { \ xla::XlaBuilder* b = ctx->builder(); \ (void)b; \ (void)lhs_shape; \ (void)rhs_shape; \ + (void)broadcast_output_exprs; \ (void)extend_dimensions; \ return HLO; \ } \ @@ -87,8 +90,12 @@ XLA_MAKE_BINARY(Complex, xla::Complex(lhs, rhs, extend_dimensions), xla::DExpr() // return x / y; // } static xla::XlaOp DivNoNanImpl(xla::XlaBuilder* b, DataType dtype, xla::XlaOp x, - xla::XlaOp y, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + xla::XlaOp y, + const absl::Span& + broadcast_output_exprs, + const BCast& broadcast_helper) { + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); auto zero = XlaHelpers::Zero(b, dtype); auto y_equals_0 = xla::Eq(y, zero); auto zeros = xla::ZerosLike(x); @@ -96,7 +103,8 @@ static xla::XlaOp DivNoNanImpl(xla::XlaBuilder* b, DataType dtype, xla::XlaOp x, return result; } XLA_MAKE_BINARY(DivNoNan, - DivNoNanImpl(b, input_type(0), lhs, rhs, broadcast_helper), + DivNoNanImpl(b, input_type(0), lhs, rhs, + broadcast_output_exprs, broadcast_helper), xla::DExpr()); // Implementation of MulNoNan. Pseudo-code: @@ -106,8 +114,12 @@ XLA_MAKE_BINARY(DivNoNan, // return x * y; // } static xla::XlaOp MulNoNanImpl(xla::XlaBuilder* b, DataType dtype, xla::XlaOp x, - xla::XlaOp y, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + xla::XlaOp y, + const absl::Span& + broadcast_output_exprs, + const BCast& broadcast_helper) { + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); auto zero = XlaHelpers::Zero(b, dtype); auto y_equals_0 = xla::Eq(y, zero); auto zeros = xla::ZerosLike(x); @@ -115,7 +127,8 @@ static xla::XlaOp MulNoNanImpl(xla::XlaBuilder* b, DataType dtype, xla::XlaOp x, return result; } XLA_MAKE_BINARY(MulNoNan, - MulNoNanImpl(b, input_type(0), lhs, rhs, broadcast_helper), + MulNoNanImpl(b, input_type(0), lhs, rhs, + broadcast_output_exprs, broadcast_helper), xla::DExpr()); // Implementation of FloorDiv. @@ -129,8 +142,12 @@ XLA_MAKE_BINARY(MulNoNan, // return z; // } static xla::XlaOp FloorDivImpl(xla::XlaBuilder* b, DataType dtype, xla::XlaOp x, - xla::XlaOp y, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + xla::XlaOp y, + const absl::Span& + broadcast_output_exprs, + const BCast& broadcast_helper) { + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); if (DataTypeIsFloating(dtype)) { if (dtype == DataType::DT_BFLOAT16) { // The result of a BF16 division may produce the Ceil of what was @@ -155,45 +172,61 @@ static xla::XlaOp FloorDivImpl(xla::XlaBuilder* b, DataType dtype, xla::XlaOp x, return xla::Select(round_down, xla::Sub(x_div_y, one), x_div_y); } XLA_MAKE_BINARY(FloorDiv, - FloorDivImpl(b, input_type(0), lhs, rhs, broadcast_helper), + FloorDivImpl(b, input_type(0), lhs, rhs, + broadcast_output_exprs, broadcast_helper), (lhs / rhs).simplify()); xla::XlaOp XlogyImpl(xla::XlaOp x, xla::XlaOp y, + const absl::Span& broadcast_output_exprs, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); auto zero = xla::ZerosLike(x); auto is_zero = xla::Eq(x, zero); return xla::Select(is_zero, zero, xla::Mul(x, xla::Log(y))); } -XLA_MAKE_BINARY(Xlogy, XlogyImpl(lhs, rhs, broadcast_helper), xla::DExpr()); +XLA_MAKE_BINARY(Xlogy, + XlogyImpl(lhs, rhs, broadcast_output_exprs, broadcast_helper), + xla::DExpr()); xla::XlaOp Xlog1pyImpl(xla::XlaOp x, xla::XlaOp y, + const absl::Span& broadcast_output_exprs, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); auto non_zero = xla::Mul(x, xla::Log1p(y)); auto zero = xla::ZerosLike(non_zero); auto x_is_zero = xla::Eq(x, zero); return xla::Select(x_is_zero, zero, non_zero); } -XLA_MAKE_BINARY(Xlog1py, Xlog1pyImpl(lhs, rhs, broadcast_helper), +XLA_MAKE_BINARY(Xlog1py, + Xlog1pyImpl(lhs, rhs, broadcast_output_exprs, broadcast_helper), xla::DExpr()); xla::XlaOp XdivyImpl(xla::XlaOp x, xla::XlaOp y, + const absl::Span& broadcast_output_exprs, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); auto zero = xla::ZerosLike(x); auto is_zero = xla::Eq(x, zero); return xla::Select(is_zero, zero, xla::Div(x, y)); } -XLA_MAKE_BINARY(Xdivy, XdivyImpl(lhs, rhs, broadcast_helper), xla::DExpr()); +XLA_MAKE_BINARY(Xdivy, + XdivyImpl(lhs, rhs, broadcast_output_exprs, broadcast_helper), + xla::DExpr()); // Implementation of FloorMod. Pseudo-code: // T trunc_mod = std::fmod(x, y); // return trunc_mod != 0 && (y < 0 != trunc_mod < 0) ? trunc_mod + y // : trunc_mod; static xla::XlaOp FloorModImpl(xla::XlaBuilder* b, DataType dtype, xla::XlaOp x, - xla::XlaOp y, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + xla::XlaOp y, + const absl::Span& + broadcast_output_exprs, + const BCast& broadcast_helper) { + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); auto zero = XlaHelpers::Zero(b, dtype); auto trunc_mod = xla::Rem(x, y); auto trunc_mod_not_zero = xla::Ne(trunc_mod, zero); @@ -202,7 +235,8 @@ static xla::XlaOp FloorModImpl(xla::XlaBuilder* b, DataType dtype, xla::XlaOp x, return xla::Select(do_plus, xla::Add(trunc_mod, y), trunc_mod); } XLA_MAKE_BINARY(FloorMod, - FloorModImpl(b, input_type(0), lhs, rhs, broadcast_helper), + FloorModImpl(b, input_type(0), lhs, rhs, + broadcast_output_exprs, broadcast_helper), xla::DExpr()); XLA_MAKE_BINARY(BitwiseAnd, xla::And(lhs, rhs, extend_dimensions), @@ -250,9 +284,13 @@ XLA_MAKE_BINARY( // For floating-point values, returns trunc(x / y). For integers, simply // returns x / y. static xla::XlaOp TruncateDivImpl(xla::XlaBuilder* b, DataType dtype, - xla::XlaOp x, xla::XlaOp y, + xla::XlaOp x, + xla::XlaOp y, + const absl::Span& + broadcast_output_exprs, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); if (!DataTypeIsFloating(dtype)) { return xla::Div(x, y); } @@ -262,7 +300,8 @@ static xla::XlaOp TruncateDivImpl(xla::XlaBuilder* b, DataType dtype, return xla::Select(round_up, xla::Ceil(x_div_y), xla::Floor(x_div_y)); } XLA_MAKE_BINARY(TruncateDiv, - TruncateDivImpl(b, input_type(0), lhs, rhs, broadcast_helper), + TruncateDivImpl(b, input_type(0), lhs, rhs, + broadcast_output_exprs, broadcast_helper), (lhs / rhs).simplify()); XLA_MAKE_BINARY(TruncateMod, xla::Rem(lhs, rhs, extend_dimensions), xla::DExpr()); @@ -318,56 +357,82 @@ XLA_MAKE_BINARY(SquaredDifference, xla::DExpr()); xla::XlaOp IgammaImpl(xla::XlaOp x, xla::XlaOp y, + const absl::Span& broadcast_output_exprs, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); return xla::Igamma(x, y); } -XLA_MAKE_BINARY(Igamma, IgammaImpl(lhs, rhs, broadcast_helper), xla::DExpr()); +XLA_MAKE_BINARY(Igamma, + IgammaImpl(lhs, rhs, broadcast_output_exprs, broadcast_helper), + xla::DExpr()); xla::XlaOp IgammaGradAImpl(xla::XlaOp x, xla::XlaOp y, + const absl::Span& + broadcast_output_exprs, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); return xla::IgammaGradA(x, y); } -XLA_MAKE_BINARY(IgammaGradA, IgammaGradAImpl(lhs, rhs, broadcast_helper), +XLA_MAKE_BINARY(IgammaGradA, + IgammaGradAImpl(lhs, rhs, broadcast_output_exprs, + broadcast_helper), xla::DExpr()); xla::XlaOp RandomGammaGradImpl(xla::XlaOp x, xla::XlaOp y, + const absl::Span& + broadcast_output_exprs, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); return xla::RandomGammaGrad(x, y); } XLA_MAKE_BINARY(RandomGammaGrad, - RandomGammaGradImpl(lhs, rhs, broadcast_helper), + RandomGammaGradImpl(lhs, rhs, broadcast_output_exprs, + broadcast_helper), xla::DExpr()); xla::XlaOp IgammacImpl(xla::XlaOp x, xla::XlaOp y, + const absl::Span& broadcast_output_exprs, const BCast& broadcast_helper) { - std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper); + std::tie(x, y) = XlaBinaryOp::Broadcast(x, y, broadcast_helper, + broadcast_output_exprs); return xla::Igammac(x, y); } -XLA_MAKE_BINARY(Igammac, IgammacImpl(lhs, rhs, broadcast_helper), +XLA_MAKE_BINARY(Igammac, + IgammacImpl(lhs, rhs, broadcast_output_exprs, broadcast_helper), xla::DExpr()); xla::XlaOp PolygammaImpl(xla::XlaOp n, xla::XlaOp x, + const absl::Span& + broadcast_output_exprs, const BCast& broadcast_helper) { - std::tie(n, x) = XlaBinaryOp::Broadcast(n, x, broadcast_helper); + std::tie(n, x) = XlaBinaryOp::Broadcast(n, x, broadcast_helper, + broadcast_output_exprs); return xla::Polygamma(n, x); } -XLA_MAKE_BINARY(Polygamma, PolygammaImpl(lhs, rhs, broadcast_helper), +XLA_MAKE_BINARY(Polygamma, + PolygammaImpl(lhs, rhs, broadcast_output_exprs, + broadcast_helper), xla::DExpr()); -xla::XlaOp ZetaImpl(xla::XlaOp x, xla::XlaOp q, const BCast& broadcast_helper) { - std::tie(x, q) = XlaBinaryOp::Broadcast(x, q, broadcast_helper); +xla::XlaOp ZetaImpl(xla::XlaOp x, xla::XlaOp q, + const absl::Span& broadcast_output_exprs, + const BCast& broadcast_helper) { + std::tie(x, q) = XlaBinaryOp::Broadcast(x, q, broadcast_helper, + broadcast_output_exprs); return xla::Zeta(x, q); } -XLA_MAKE_BINARY(Zeta, ZetaImpl(lhs, rhs, broadcast_helper), xla::DExpr()); +XLA_MAKE_BINARY(Zeta, + ZetaImpl(lhs, rhs, broadcast_output_exprs, broadcast_helper), + xla::DExpr()); #undef XLA_MAKE_BINARY 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/const_op.cc b/tensorflow/compiler/tf2xla/kernels/const_op.cc index a911f28246da77..33d8b2079897d0 100644 --- a/tensorflow/compiler/tf2xla/kernels/const_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/const_op.cc @@ -35,6 +35,9 @@ limitations under the License. namespace tensorflow { namespace { +constexpr char kUserInferredValueContentsAttrName[] = + "_user_inferred_value_contents"; + template DstT CastTo(SrcT src) { return static_cast(src); @@ -107,61 +110,6 @@ xla::XlaOp GetScalarConst(const TensorProto& proto, xla::XlaBuilder* b) { return xla::XlaOp(); } -bool IsDynamicExpressionProto(const ExpressionProto& proto) { - switch (proto.node_type_case()) { - case ExpressionProto::kVariableId: - return true; - case ExpressionProto::kAddNode: - return IsDynamicExpressionProto(proto.add_node().lhs()) || - IsDynamicExpressionProto(proto.add_node().rhs()); - case ExpressionProto::kSubNode: - return IsDynamicExpressionProto(proto.sub_node().lhs()) || - IsDynamicExpressionProto(proto.sub_node().rhs()); - case ExpressionProto::kMulNode: - return IsDynamicExpressionProto(proto.mul_node().lhs()) || - IsDynamicExpressionProto(proto.mul_node().rhs()); - case ExpressionProto::kDivNode: - return IsDynamicExpressionProto(proto.div_node().lhs()) || - IsDynamicExpressionProto(proto.div_node().rhs()); - case ExpressionProto::kConstantValue: - case ExpressionProto::NODE_TYPE_NOT_SET: - return false; - } -} - -static xla::DExpr DimExprToDExpr(const DimExpr* e) { - if (e == nullptr) { - return xla::DExpr(); - } - switch (e->kind()) { - case DimExpr::Kind::kConstant: { - const auto* ac = static_cast(e); - return xla::DExpr::Const(ac->value()); - } - case DimExpr::Kind::kVariable: { - const auto* av = static_cast(e); - return xla::DExpr::Var(av->id()); - } - case DimExpr::Kind::kAdd: { - const auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) + DimExprToDExpr(ee->rhs()); - } - case DimExpr::Kind::kSub: { - const auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) - DimExprToDExpr(ee->rhs()); - } - case DimExpr::Kind::kMul: { - const auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) * DimExprToDExpr(ee->rhs()); - } - case DimExpr::Kind::kDiv: { - const auto* ee = static_cast(e); - return DimExprToDExpr(ee->lhs()) / DimExprToDExpr(ee->rhs()); - } - } - return xla::DExpr(); -} - std::vector BuildShapeContentsFromTensorShapeProto( const TensorShapeProto& shape) { std::vector contents; @@ -169,8 +117,7 @@ std::vector BuildShapeContentsFromTensorShapeProto( for (int i = 0; i < shape.dim_size(); ++i) { xla::DExpr expr; if (i < shape.expressions_size()) { - auto tf_expr = DimExpr::FromProto(shape.expressions(i)); - expr = DimExprToDExpr(tf_expr.get()); + expr = DimExprFromProto(shape.expressions(i)); } if (expr) { xla::StringPrinter printer; @@ -190,14 +137,11 @@ 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) || + (tensor_shape.dims() == 1 && + tensor_shape.dim_size(0) == contents.dim_size()); } class ConstOp : public XlaOpKernel { @@ -219,17 +163,25 @@ 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) { - 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); + GetNodeAttr(ctx->op_kernel().def(), "user_inferred_shape", + &inferred_shape_proto) + .IgnoreError(); + } + 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(); } } + 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. @@ -246,14 +198,10 @@ 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()) { - VLOG(1) << "ConstOp attaching shape contents through broadcast fast " - << "path with " << shape.dim_size(0) - << " entries and dynamic_exprs=" - << CountDynamicShapeContents(inferred_shape_proto); + if ((has_contents_proto || has_dynamic) && + CanAttachContentsFromTensorShapeProto(shape, contents_proto)) { output.set_contents( - BuildShapeContentsFromTensorShapeProto(inferred_shape_proto)); + BuildShapeContentsFromTensorShapeProto(contents_proto)); } ctx->SetOutputExpression(0, output); return; @@ -264,19 +212,11 @@ class ConstOp : public XlaOpKernel { OP_REQUIRES(ctx, tensor.FromProto(cpu_allocator(), proto_), errors::InvalidArgument("Cannot parse tensor from proto: ", proto_.DebugString())); - if (has_dynamic) { - VLOG(1) << "ConstOp tensor path tensor_shape=" - << tensor.shape().DebugString() << " inferred_rank=" - << inferred_shape_proto.dim_size(); - } XlaExpression output = XlaExpression::Constant(tensor); - if (has_dynamic && tensor.dims() == 1 && - tensor.dim_size(0) == inferred_shape_proto.dim_size()) { - VLOG(1) << "ConstOp attaching shape contents to folded const with " - << tensor.dim_size(0) << " entries and dynamic_exprs=" - << CountDynamicShapeContents(inferred_shape_proto); + if ((has_contents_proto || has_dynamic) && + CanAttachContentsFromTensorShapeProto(tensor.shape(), contents_proto)) { output.set_contents( - BuildShapeContentsFromTensorShapeProto(inferred_shape_proto)); + BuildShapeContentsFromTensorShapeProto(contents_proto)); } ctx->SetOutputExpression(0, output); } diff --git a/tensorflow/compiler/tf2xla/kernels/cwise_ops.cc b/tensorflow/compiler/tf2xla/kernels/cwise_ops.cc index 05c10e6957826e..a796509e36a4ba 100644 --- a/tensorflow/compiler/tf2xla/kernels/cwise_ops.cc +++ b/tensorflow/compiler/tf2xla/kernels/cwise_ops.cc @@ -240,6 +240,7 @@ void XlaBinaryOp::Compile(XlaOpKernelContext* ctx) { lhs = xla::SliceInDim(lhs, 0, rhs_xla_shape.dimensions(i), 1, /*dimno=*/i); lhs_tensor_shape->set_dim(i, rhs_xla_shape.dimensions(i)); + lhs_tensor_shape->set_expression(i, rhs_xla_shape.expressions(i)); // Propagate dynamic dimension. lhs = xla::SetDimensionSize(lhs, size, i); } @@ -262,6 +263,7 @@ void XlaBinaryOp::Compile(XlaOpKernelContext* ctx) { lhs, xla::Zero(ctx->builder(), lhs_xla_shape.element_type()), i, 0, diff); lhs_tensor_shape->set_dim(i, rhs_xla_shape.dimensions(i)); + lhs_tensor_shape->set_expression(i, rhs_xla_shape.expressions(i)); // Propagate dynamic dimension. lhs = xla::SetDimensionSize(lhs, size, i); } @@ -309,6 +311,10 @@ void XlaBinaryOp::Compile(XlaOpKernelContext* ctx) { lhs = xla::SetDimensionSize(lhs, size, i); lhs_tensor_shape->set_dim(i, rhs_xla_shape.dimensions(i)); + lhs_tensor_shape->set_expression( + i, (lhs_tensor_shape->get_filled_expression(i) * + rhs_xla_shape.expressions(i)) + .simplify()); } } } @@ -334,6 +340,67 @@ void XlaBinaryOp::Compile(XlaOpKernelContext* ctx) { return; } + auto build_broadcast_output_expressions = + [&lhs_shape, &rhs_shape, &bcast]() -> std::vector { + auto merge_broadcast_dim = [](bool has_lhs, int64_t lhs_dim, + const xla::DExpr& lhs_expr, bool has_rhs, + int64_t rhs_dim, const xla::DExpr& rhs_expr, + int64_t output_dim) { + if (!has_lhs) { + return rhs_expr; + } + if (!has_rhs) { + return lhs_expr; + } + if (lhs_dim == 1 && rhs_dim != 1) { + // A broadcasted singleton usually inherits the other side's + // expression, but keep a dynamic singleton visible by folding it into + // the result. + return lhs_expr && lhs_expr->is_dynamic() + ? (lhs_expr * rhs_expr).simplify() + : rhs_expr; + } + if (rhs_dim == 1 && lhs_dim != 1) { + // Symmetric case for a singleton rhs broadcast. + return rhs_expr && rhs_expr->is_dynamic() + ? (rhs_expr * lhs_expr).simplify() + : lhs_expr; + } + // When both sides describe the same logical dimension, prefer whichever + // side still carries a dynamic symbolic expression. + if (lhs_expr && lhs_expr->is_dynamic()) { + return lhs_expr; + } + if (rhs_expr && rhs_expr->is_dynamic()) { + return rhs_expr; + } + return xla::DExpr::Const(output_dim); + }; + + const auto& output_shape = bcast.output_shape(); + std::vector output_exprs(output_shape.size()); + + for (int out_i = output_shape.size() - 1, lhs_i = lhs_shape.dims() - 1, + rhs_i = rhs_shape.dims() - 1; + out_i >= 0; --out_i, --lhs_i, --rhs_i) { + const bool has_lhs = lhs_i >= 0; + const bool has_rhs = rhs_i >= 0; + xla::DExpr lhs_expr = has_lhs ? lhs_shape.get_filled_expression(lhs_i) + : xla::DExpr::Const(1); + xla::DExpr rhs_expr = has_rhs ? rhs_shape.get_filled_expression(rhs_i) + : xla::DExpr::Const(1); + const int64_t lhs_dim = has_lhs ? lhs_shape.dim_size(lhs_i) : 1; + const int64_t rhs_dim = has_rhs ? rhs_shape.dim_size(rhs_i) : 1; + output_exprs[out_i] = merge_broadcast_dim( + has_lhs, lhs_dim, lhs_expr, has_rhs, rhs_dim, rhs_expr, + output_shape[out_i]); + } + + return output_exprs; + }; + std::vector broadcast_output_exprs = + build_broadcast_output_expressions(); + // If the ranks of the inputs don't match, TensorFlow automatically // reshapes the smaller by padding with dimensions of size 1 as a // prefix. In other words to pad a 5-vector to a 3-dimensional @@ -359,7 +426,8 @@ void XlaBinaryOp::Compile(XlaOpKernelContext* ctx) { // Call virtual method to emit the computation. xla::XlaOp output = Computation(ctx, lhs_handle, lhs_shape.dim_sizes(), rhs_handle, - rhs_shape.dim_sizes(), bcast, extend_dimension); + rhs_shape.dim_sizes(), bcast, + broadcast_output_exprs, extend_dimension); // The TensorFlow helper computed the post-broadcast shape in // output_shape: we rely on subclassed Computations to implement the @@ -377,13 +445,20 @@ void XlaBinaryOp::Compile(XlaOpKernelContext* ctx) { } /* static */ std::pair XlaBinaryOp::Broadcast( - xla::XlaOp lhs, xla::XlaOp rhs, const BCast& broadcast_helper) { - auto lhs_output = BroadcastTo(lhs, broadcast_helper.output_shape()); + xla::XlaOp lhs, xla::XlaOp rhs, const BCast& broadcast_helper, + absl::Span output_exprs) { + CHECK_EQ(output_exprs.size(), broadcast_helper.output_shape().size()); + for (const xla::DExpr& expr : output_exprs) { + CHECK(expr); + } + auto lhs_output = + BroadcastTo(lhs, broadcast_helper.output_shape(), output_exprs); if (!lhs_output.ok()) { xla::XlaOp error = lhs.builder()->ReportError(lhs_output.status()); return {error, error}; } - auto rhs_output = BroadcastTo(rhs, broadcast_helper.output_shape()); + auto rhs_output = + BroadcastTo(rhs, broadcast_helper.output_shape(), output_exprs); if (!rhs_output.ok()) { xla::XlaOp error = rhs.builder()->ReportError(rhs_output.status()); return {error, error}; diff --git a/tensorflow/compiler/tf2xla/kernels/cwise_ops.h b/tensorflow/compiler/tf2xla/kernels/cwise_ops.h index df92dac718db13..6e3734abe7d53f 100644 --- a/tensorflow/compiler/tf2xla/kernels/cwise_ops.h +++ b/tensorflow/compiler/tf2xla/kernels/cwise_ops.h @@ -65,6 +65,7 @@ class XlaBinaryOp : public XlaOpKernel { XlaOpKernelContext* ctx, const xla::XlaOp& lhs, const absl::Span& lhs_shape, const xla::XlaOp& rhs, const absl::Span& rhs_shape, const BCast& broadcast_helper, + const absl::Span& broadcast_output_exprs, const std::vector& extend_dimensions) = 0; // Returns a symbolic expression for one output element when content metadata @@ -81,7 +82,8 @@ class XlaBinaryOp : public XlaOpKernel { // 'broadcast_helper', yielding arguments 'lhs' and 'rhs' that have the same // shape. static std::pair Broadcast( - xla::XlaOp lhs, xla::XlaOp rhs, const BCast& broadcast_helper); + xla::XlaOp lhs, xla::XlaOp rhs, const BCast& broadcast_helper, + absl::Span output_exprs); }; } // namespace tensorflow 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..46404dc31ed23f 100644 --- a/tensorflow/compiler/tf2xla/kernels/dynamic_partition_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/dynamic_partition_op.cc @@ -104,7 +104,7 @@ class DynamicPartitionOp : public XlaOpKernel { xla::XlaOp valid_element = xla::Lt(input_index, dynamic_input_count); xla::XlaOp invalid_partition = xla::Broadcast(xla::ConstantR0(ctx->builder(), num_partitions_), - {input_count}); + {input_count}, partition_1d_shape.expressions()); partitions_1d = xla::Select(valid_element, partitions_1d, invalid_partition); std::vector to_sort = {partitions_1d, data_1d}; @@ -213,11 +213,12 @@ 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{CollapseExpressions(data_exprs)}); xla::Shape partitions_1d_shape = xla::ShapeUtil::MakeShape( partition_shape.element_type(), {input_count}, - {xla::DExpr::Const(input_count)}); + std::vector{ + CollapseExpressions(flattened_partition_exprs)}); 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/fft_ops.cc b/tensorflow/compiler/tf2xla/kernels/fft_ops.cc index 8fb04773aafb49..5d4114d8c88d5d 100644 --- a/tensorflow/compiler/tf2xla/kernels/fft_ops.cc +++ b/tensorflow/compiler/tf2xla/kernels/fft_ops.cc @@ -20,8 +20,8 @@ limitations under the License. #include #include "absl/container/inlined_vector.h" -#include "tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h" #include "tensorflow/compiler/tf2xla/xla_helpers.h" +#include "tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h" #include "tensorflow/compiler/tf2xla/xla_op_kernel.h" #include "tensorflow/compiler/tf2xla/xla_op_registry.h" #include "xla/hlo/builder/xla_builder.h" @@ -139,9 +139,9 @@ class IFFTOp : public GenericFftOp { explicit IFFTOp(OpKernelConstruction* ctx) : GenericFftOp(ctx, /*fft_type=*/FftType::IFFT, /*fft_rank=*/FFTRank) {} }; -REGISTER_XLA_OP(Name("IFFT").TypeConstraint("Tcomplex", - {DT_COMPLEX64, DT_COMPLEX128}), - MlirXlaOpKernel); +REGISTER_XLA_OP_FACTORY( + Name("IFFT").TypeConstraint("Tcomplex", {DT_COMPLEX64, DT_COMPLEX128}), + CreateDynamicNativeXlaOpKernel>); REGISTER_XLA_OP(Name("IFFT2D").TypeConstraint("Tcomplex", {DT_COMPLEX64, DT_COMPLEX128}), IFFTOp<2>); diff --git a/tensorflow/compiler/tf2xla/kernels/identity_op.cc b/tensorflow/compiler/tf2xla/kernels/identity_op.cc index 7b04d501d305de..f6abb610c456e4 100644 --- a/tensorflow/compiler/tf2xla/kernels/identity_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/identity_op.cc @@ -62,7 +62,8 @@ REGISTER_XLA_OP(Name("IdentityN") .CompilationOnly(), IdentityOp); REGISTER_XLA_OP(Name("PlaceholderWithDefault"), IdentityOp); -REGISTER_XLA_OP(Name("PreventGradient"), MlirXlaOpKernel); +REGISTER_XLA_OP_FACTORY( + Name("PreventGradient"), CreateDynamicNativeXlaOpKernel); REGISTER_XLA_OP(Name("StopGradient").AllowVariantTypes(), IdentityOp); REGISTER_XLA_OP(Name("Snapshot"), IdentityOp); REGISTER_XLA_OP(Name("_EagerConst"), IdentityOp); diff --git a/tensorflow/compiler/tf2xla/kernels/resampler_ops.cc b/tensorflow/compiler/tf2xla/kernels/resampler_ops.cc index c54c4613d29e44..602d480b5b6541 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_filled_expressions(); // 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,17 @@ 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_filled_expressions(); // 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.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/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/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..cd011016a4672c 100644 --- a/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc +++ b/tensorflow/compiler/tf2xla/kernels/sequence_ops.cc @@ -51,12 +51,6 @@ xla::DExpr GetScalarExpr(const XlaExpression& expression, return xla::DExpr::Const(literal.Get({})); } -bool HasStaticScalarContent(const XlaExpression& expression) { - const auto& contents = expression.contents(); - return contents.empty() || - (contents[0] && contents[0]->is_constant()); -} - bool HasDynamicContent(const XlaExpression& expression) { return absl::c_any_of(expression.contents(), [](const xla::DExpr& expr) { return expr && expr->is_dynamic(); @@ -92,17 +86,30 @@ 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(); + const auto& start_contents = start_expr.contents(); + xla::DExpr effective_start = + (!start_contents.empty() && start_contents[0]) ? start_contents[0] + : start_symbol; + const auto& limit_contents = limit_expr.contents(); + xla::DExpr effective_limit = + (!limit_contents.empty() && limit_contents[0]) ? limit_contents[0] + : limit_symbol; + const auto& delta_contents = delta_expr.contents(); + xla::DExpr effective_delta = + (!delta_contents.empty() && delta_contents[0]) ? delta_contents[0] + : delta_symbol; + + xla::DExpr positive_diff = (effective_limit - effective_start).simplify(); + xla::DExpr positive_size = + (((positive_diff - 1) / effective_delta) + 1).simplify(); + xla::DExpr negative_step = + (xla::DExpr::Const(0) - effective_delta).simplify(); + xla::DExpr negative_diff = (effective_start - effective_limit).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. @@ -143,7 +150,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); @@ -188,13 +195,14 @@ class RangeOp : public XlaOpKernel { : (std::abs(limit_value - start_value) - 1) / std::abs(delta_value) + 1); - xla::DExpr size_expr = - HasStaticScalarContent(ctx->InputExpression(2)) - ? BuildRangeSizeExpr(ctx->InputExpression(0), - ctx->InputExpression(1), - ctx->InputExpression(2), start, - limit, delta, size) - : xla::DExpr::Const(size); + xla::DExpr size_expr = xla::DExpr::Const(size); + if (HasDynamicContent(ctx->InputExpression(0)) || + HasDynamicContent(ctx->InputExpression(1)) || + HasDynamicContent(ctx->InputExpression(2))) { + size_expr = BuildRangeSizeExpr( + ctx->InputExpression(0), ctx->InputExpression(1), + ctx->InputExpression(2), start, limit, delta, size); + } output = CreateRangeTensor(start, limit, delta, ctx->builder(), size_expr); break; @@ -209,13 +217,14 @@ class RangeOp : public XlaOpKernel { : (std::abs(limit_value - start_value) - 1) / std::abs(delta_value) + 1; - xla::DExpr size_expr = - HasStaticScalarContent(ctx->InputExpression(2)) - ? BuildRangeSizeExpr(ctx->InputExpression(0), - ctx->InputExpression(1), - ctx->InputExpression(2), start, - limit, delta, size) - : xla::DExpr::Const(size); + xla::DExpr size_expr = xla::DExpr::Const(size); + if (HasDynamicContent(ctx->InputExpression(0)) || + HasDynamicContent(ctx->InputExpression(1)) || + HasDynamicContent(ctx->InputExpression(2))) { + size_expr = BuildRangeSizeExpr( + ctx->InputExpression(0), ctx->InputExpression(1), + ctx->InputExpression(2), start, limit, delta, size); + } output = CreateRangeTensor(start, limit, delta, ctx->builder(), size_expr); break; @@ -256,10 +265,12 @@ class RangeOp : public XlaOpKernel { } const XlaExpression& start_expr = ctx->InputExpression(0); + const XlaExpression& limit_expr = ctx->InputExpression(1); const XlaExpression& delta_expr = ctx->InputExpression(2); const bool symbolic_enabled = SymbolicContentEnabled(); const bool has_dynamic_content = - HasDynamicContent(start_expr) || HasDynamicContent(delta_expr); + HasDynamicContent(start_expr) || HasDynamicContent(limit_expr) || + HasDynamicContent(delta_expr); if (type == DT_INT32) { int32 start_value = start.Get({}); diff --git a/tensorflow/compiler/tf2xla/kernels/shape_op.cc b/tensorflow/compiler/tf2xla/kernels/shape_op.cc index ed5c1c24400e0d..c47cad63e01dae 100644 --- a/tensorflow/compiler/tf2xla/kernels/shape_op.cc +++ b/tensorflow/compiler/tf2xla/kernels/shape_op.cc @@ -314,6 +314,7 @@ class SizeOp : public XlaOpKernel { const TensorShape input_shape = ctx->InputShape(0); xla::XlaBuilder* builder = ctx->builder(); auto size = xla::One(builder, ctx->output_xla_type(0)); + xla::DExpr size_expr = xla::DExpr::Const(1); const int rank = input_shape.dims(); for (int64_t dim = 0; dim < rank; ++dim) { @@ -326,12 +327,39 @@ class SizeOp : public XlaOpKernel { "on all dimensions, found ", input_shape.dim_size(dim), " elements on dimension ", dim))); - size = xla::Mul(size, xla::ConvertElementType( - xla::GetDimensionSize(ctx->Input(0), dim), - ctx->output_xla_type(0))); + xla::DExpr dim_expr = input_shape.get_filled_expression(dim); + std::vector contents = { + dim_expr && dim_expr->is_dynamic() + ? dim_expr + : xla::DExpr::Unknown(xla::kUnknownContentSentinel)}; + xla::XlaOp dim_size = xla::GetDimensionSize(ctx->Input(0), dim); + if (SymbolicContentEnabled()) { + OP_REQUIRES_OK(ctx, + builder->SetInstructionContents(dim_size, contents)); + } + xla::XlaOp converted = + xla::ConvertElementType(dim_size, ctx->output_xla_type(0)); + if (SymbolicContentEnabled()) { + OP_REQUIRES_OK(ctx, + builder->SetInstructionContents(converted, contents)); + } + size = xla::Mul(size, converted); + size_expr = + (size_expr * input_shape.get_filled_expression(dim)).simplify(); + if (SymbolicContentEnabled()) { + OP_REQUIRES_OK(ctx, + builder->SetInstructionContents(size, {size_expr})); + } } - ctx->SetOutput(0, size); + if (SymbolicContentEnabled()) { + XlaExpression output = + XlaExpression::XlaOp(size, ctx->expected_output_dtype(0)); + output.set_contents({size_expr}); + ctx->SetOutputExpression(0, output); + } else { + ctx->SetOutput(0, 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..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); } @@ -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/unary_ops.cc b/tensorflow/compiler/tf2xla/kernels/unary_ops.cc index f35e375356516b..8dc75be96d7958 100644 --- a/tensorflow/compiler/tf2xla/kernels/unary_ops.cc +++ b/tensorflow/compiler/tf2xla/kernels/unary_ops.cc @@ -45,6 +45,22 @@ namespace { }; \ REGISTER_XLA_OP(Name(#NAME), NAME##Op); +#define XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(NAME, COMPUTATION) \ + class NAME##NativeOp : public XlaOpKernel { \ + public: \ + explicit NAME##NativeOp(OpKernelConstruction* ctx) \ + : XlaOpKernel(ctx) {} \ + void Compile(XlaOpKernelContext* ctx) override { \ + xla::XlaBuilder* b = ctx->builder(); \ + (void)b; \ + xla::XlaOp x = ctx->Input(0); \ + xla::XlaOp y = COMPUTATION; \ + ctx->SetOutput(0, y); \ + } \ + }; \ + REGISTER_XLA_OP_FACTORY( \ + Name(#NAME), CreateDynamicNativeXlaOpKernel); + XLAJIT_MAKE_UNARY(ComplexAbs, xla::Abs(x)); XLAJIT_MAKE_UNARY(Angle, xla::Atan2(xla::Imag(x), xla::Real(x))); @@ -52,29 +68,30 @@ XLAJIT_MAKE_UNARY(Angle, xla::Atan2(xla::Imag(x), xla::Real(x))); XLAJIT_MAKE_UNARY(Conj, xla::Conj(x)); // Return x if x>0, otherwise -x. -REGISTER_XLA_OP(Name("Abs"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Abs, xla::Abs(x)); XLAJIT_MAKE_UNARY(Acos, xla::Acos(x)); XLAJIT_MAKE_UNARY(Acosh, xla::Acosh(x)); XLAJIT_MAKE_UNARY(Asin, xla::Asin(x)) XLAJIT_MAKE_UNARY(Asinh, xla::Asinh(x)); -REGISTER_XLA_OP(Name("Atan"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Atan, xla::Atan(x)); XLAJIT_MAKE_UNARY(Atanh, xla::Atanh(x)); -REGISTER_XLA_OP(Name("Ceil"), MlirXlaOpKernel); -REGISTER_XLA_OP(Name("Cos"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Ceil, xla::Ceil(x)); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Cos, xla::Cos(x)); XLAJIT_MAKE_UNARY(Cosh, xla::Cosh(x)); XLAJIT_MAKE_UNARY(Sin, xla::Sin(x)); XLAJIT_MAKE_UNARY(Tan, xla::Tan(x)); -REGISTER_XLA_OP(Name("Exp"), MlirXlaOpKernel); -REGISTER_XLA_OP(Name("Expm1"), MlirXlaOpKernel); -REGISTER_XLA_OP(Name("Floor"), MlirXlaOpKernel); -REGISTER_XLA_OP(Name("IsFinite"), MlirXlaOpKernel); -REGISTER_XLA_OP(Name("IsInf"), MlirXlaOpKernel); -REGISTER_XLA_OP(Name("IsNan"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Exp, xla::Exp(x)); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Expm1, xla::Expm1(x)); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Floor, xla::Floor(x)); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(IsFinite, xla::IsFinite(x)); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(IsInf, xla::IsInf(x)); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(IsNan, xla::IsNan(x)); // Return 1/x XLAJIT_MAKE_UNARY(Inv, xla::ScalarLike(x, 1.0) / x); -REGISTER_XLA_OP(Name("Reciprocal"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK( + Reciprocal, xla::ScalarLike(x, 1.0) / x); XLAJIT_MAKE_UNARY(Log, xla::Log(x)); -REGISTER_XLA_OP(Name("Log1p"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Log1p, xla::Log1p(x)); XLAJIT_MAKE_UNARY(Invert, xla::Not(x)); XLAJIT_MAKE_UNARY(LogicalNot, xla::Not(x)); @@ -85,12 +102,12 @@ XLAJIT_MAKE_UNARY(Neg, -x); XLAJIT_MAKE_UNARY(Rint, xla::RoundToEven(x)); XLAJIT_MAKE_UNARY(Round, xla::RoundToEven(x)); -REGISTER_XLA_OP(Name("Rsqrt"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Rsqrt, xla::Rsqrt(x)); -REGISTER_XLA_OP(Name("Sigmoid"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Sigmoid, xla::Logistic(x)); // Returns NaN if x is NaN, 0 if x is 0, -1 if x < 0 and 1 if x > 0. -REGISTER_XLA_OP(Name("Sign"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Sign, xla::Sign(x)); XLAJIT_MAKE_UNARY(Sinh, xla::Sinh(x)); static xla::XlaOp Softplus(xla::XlaBuilder* b, xla::XlaOp features) { @@ -115,11 +132,11 @@ XLAJIT_MAKE_UNARY(Softplus, Softplus(b, x)); // softsign(x) = x / (abs(x) + 1) XLAJIT_MAKE_UNARY(Softsign, x / (xla::Abs(x) + xla::ScalarLike(x, 1.0))); -REGISTER_XLA_OP(Name("Sqrt"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Sqrt, xla::Sqrt(x)); XLAJIT_MAKE_UNARY(Square, x* x); -REGISTER_XLA_OP(Name("Tanh"), MlirXlaOpKernel); -REGISTER_XLA_OP(Name("Real"), MlirXlaOpKernel); -REGISTER_XLA_OP(Name("Imag"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Tanh, xla::Tanh(x)); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Real, xla::Real(x)); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Imag, xla::Imag(x)); XLAJIT_MAKE_UNARY(Erf, xla::Erf(x)); XLAJIT_MAKE_UNARY(Erfc, xla::Erfc(x)); XLAJIT_MAKE_UNARY(Erfinv, xla::ErfInv(x)); @@ -127,7 +144,7 @@ XLAJIT_MAKE_UNARY(Erfinv, xla::ErfInv(x)); XLAJIT_MAKE_UNARY(Ndtri, xla::ScalarLike(x, std::sqrt(2.0)) * xla::ErfInv(xla::ScalarLike(x, 2.0) * x - xla::ScalarLike(x, 1.0))); -REGISTER_XLA_OP(Name("Lgamma"), MlirXlaOpKernel); +XLAJIT_MAKE_UNARY_WITH_MLIR_FALLBACK(Lgamma, xla::Lgamma(x)); XLAJIT_MAKE_UNARY(Digamma, xla::Digamma(x)); XLAJIT_MAKE_UNARY(BesselI0e, xla::BesselI0e(x)); XLAJIT_MAKE_UNARY(BesselI1e, xla::BesselI1e(x)); 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/mlir_xla_op_kernel.cc b/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.cc index b1a93508d92896..ca80ded3aeb257 100644 --- a/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.cc +++ b/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.cc @@ -15,16 +15,21 @@ limitations under the License. #include "tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h" +#include #include +#include +#include #include "absl/status/status.h" #include "absl/strings/str_cat.h" #include "llvm/ADT/DenseSet.h" #include "llvm/ADT/SmallVector.h" #include "mlir/IR/MLIRContext.h" // from @llvm-project +#include "tensorflow/compiler/jit/shape_inference.h" #include "tensorflow/compiler/jit/xla_compile_util.h" #include "tensorflow/compiler/mlir/tf2xla/api/v1/compile_mlir_util.h" #include "tensorflow/compiler/mlir/utils/array_container_utils.h" +#include "tensorflow/compiler/tf2xla/shape_util.h" #include "tensorflow/compiler/tf2xla/xla_compiler.h" #include "tensorflow/compiler/tf2xla/xla_expression.h" #include "tensorflow/compiler/tf2xla/xla_op_kernel.h" @@ -35,7 +40,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 +75,55 @@ class MLIRContextResource : public ResourceBase { mlir::MLIRContext mlir_ctx_; }; +bool HasDynamicExpressions(const xla::Shape& shape) { + for (int64_t dim = 0; dim < shape.dimensions_size(); ++dim) { + const xla::DExpr& expression = shape.expressions(dim); + if (expression && expression->is_dynamic()) return true; + } + return false; +} + +absl::StatusOr StripDynamicExpressions(xla::XlaOp input) { + TF_ASSIGN_OR_RETURN(xla::Shape shape, input.builder()->GetShape(input)); + if (!HasDynamicExpressions(shape)) return input; + + xla::Shape static_shape = shape; + for (int64_t dim = 0; dim < shape.dimensions_size(); ++dim) { + static_shape.set_expression(dim, xla::DExpr::Const(shape.dimensions(dim))); + } + xla::XlaOp result = xla::Reshape(static_shape, input); + TF_RETURN_IF_ERROR(input.builder()->GetShape(result).status()); + return result; +} + +absl::StatusOr RestoreOutputExpressions( + xla::XlaOp output, const PartialTensorShape& inferred_shape) { + TF_ASSIGN_OR_RETURN(xla::Shape output_shape, + output.builder()->GetShape(output)); + if (inferred_shape.unknown_rank() || + inferred_shape.dims() != output_shape.dimensions_size()) { + return errors::InvalidArgument( + "MLIR output rank does not match its TensorFlow-inferred shape"); + } + + bool has_dynamic_expression = false; + for (int64_t dim = 0; dim < output_shape.dimensions_size(); ++dim) { + const xla::DExpr& expression = inferred_shape.get_expression(dim); + if (expression && expression->is_dynamic()) { + output_shape.set_expression(dim, expression); + has_dynamic_expression = true; + } else { + output_shape.set_expression( + dim, xla::DExpr::Const(output_shape.dimensions(dim))); + } + } + if (!has_dynamic_expression) return output; + + xla::XlaOp result = xla::Reshape(output_shape, output); + TF_RETURN_IF_ERROR(output.builder()->GetShape(result).status()); + return result; +} + } // namespace absl::Status MlirXlaOpKernel::ContextToXlaArgs( @@ -144,6 +200,46 @@ absl::Status MlirXlaOpKernel::ConstructXlaOp(XlaOpKernelContext* ctx) { TF_ASSIGN_OR_RETURN(auto graph, CreateSingleOpGraph(def(), xla_args, result_dtypes)); + bool has_dynamic_expressions = false; + std::map input_shapes; + for (int i = 0; i < xla_args.size(); ++i) { + if (!std::holds_alternative(xla_args[i].shape)) continue; + const xla::Shape& xla_shape = std::get(xla_args[i].shape); + has_dynamic_expressions |= HasDynamicExpressions(xla_shape); + TensorShape tensor_shape; + TF_RETURN_IF_ERROR(XLAShapeToTensorShape(xla_shape, &tensor_shape)); + input_shapes.emplace(i, InferredShape{PartialTensorShape(tensor_shape)}); + } + + std::vector inferred_output_shapes; + if (has_dynamic_expressions) { + // Infer the logical output expressions before hiding the input expressions + // from the MLIR bridge. MLIR then lowers an ordinary statically shaped + // region, and the inferred expressions are restored at its boundary. + GraphShapeInfo shape_info; + TF_RETURN_IF_ERROR(InferShapes( + graph.get(), input_shapes, + ctx->function_library()->GetFunctionLibraryDefinition(), &shape_info)); + auto output_shapes = shape_info.find(def().name()); + if (output_shapes == shape_info.end() || + output_shapes->second.size() != result_dtypes.size()) { + return errors::InvalidArgument( + "TensorFlow shape inference did not produce all MLIR op outputs"); + } + inferred_output_shapes.reserve(output_shapes->second.size()); + for (const InferredShape& output_shape : output_shapes->second) { + inferred_output_shapes.push_back(output_shape.shape); + } + + for (int i = 0; i < xla_params.size(); ++i) { + if (xla_args[i].kind != XlaCompiler::Argument::kParameter) continue; + TF_ASSIGN_OR_RETURN(xla_params[i], + StripDynamicExpressions(xla_params[i])); + TF_ASSIGN_OR_RETURN(xla_args[i].shape, + xla_params[i].builder()->GetShape(xla_params[i])); + } + } + ResourceMgr* res_manager = ctx->op_kernel_context()->resource_manager(); MLIRContextResource* ctx_res; TF_RETURN_IF_ERROR(res_manager->LookupOrCreate( @@ -174,6 +270,11 @@ absl::Status MlirXlaOpKernel::ConstructXlaOp(XlaOpKernelContext* ctx) { // Set context outputs. for (int i = 0, end = returns.size(); i < end; ++i) { + if (has_dynamic_expressions) { + TF_ASSIGN_OR_RETURN( + returns[i], + RestoreOutputExpressions(returns[i], inferred_output_shapes[i])); + } ctx->SetOutput(i, returns[i]); } diff --git a/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h b/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h index 6053f5d68635d0..7b70096851b3df 100644 --- a/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h +++ b/tensorflow/compiler/tf2xla/mlir_xla_op_kernel.h @@ -16,6 +16,7 @@ limitations under the License. #ifndef TENSORFLOW_COMPILER_TF2XLA_MLIR_XLA_OP_KERNEL_H_ #define TENSORFLOW_COMPILER_TF2XLA_MLIR_XLA_OP_KERNEL_H_ +#include "tensorflow/compiler/jit/flags.h" #include "tensorflow/compiler/tf2xla/xla_compiler.h" #include "tensorflow/compiler/tf2xla/xla_op_kernel.h" #include "tensorflow/core/framework/op_kernel.h" @@ -36,6 +37,14 @@ class MlirXlaOpKernel : public XlaOpKernel { absl::Status ConstructXlaOp(XlaOpKernelContext* ctx); }; +template +OpKernel* CreateDynamicNativeXlaOpKernel(OpKernelConstruction* ctx) { + if (GetMarkForCompilationPassFlags()->tf_xla_enable_dynamic_sizes) { + return new NativeOp(ctx); + } + return new MlirXlaOpKernel(ctx); +} + } // namespace tensorflow #endif // TENSORFLOW_COMPILER_TF2XLA_MLIR_XLA_OP_KERNEL_H_ 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..56cb96b7a5a15c 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" @@ -36,6 +38,7 @@ limitations under the License. #include "xla/client/client_library.h" #include "xla/client/local_client.h" #include "xla/hlo/builder/xla_builder.h" +#include "xla/hlo/ir/hlo_opcode.h" #include "xla/literal.h" #include "xla/service/hlo.pb.h" #include "xla/service/hlo_module_util.h" @@ -96,6 +99,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 +139,2766 @@ 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))); +} + +// MLIR-only kernels lower against static physical shapes, then return to the +// surrounding XLA computation with the inferred expression intact. +TEST_F(XlaCompilerDynamicSizesTest, MlirKernelPreservesExpressions) { + Scope scope = Scope::NewRootScope().ExitOnError(); + auto gradients = ops::_Arg(scope.WithOpName("gradients"), DT_FLOAT, 0); + auto features = ops::_Arg(scope.WithOpName("features"), DT_FLOAT, 1); + Node* relu_grad; + auto builder = NodeBuilder("relu_grad", "ReluGrad") + .Input(ops::AsNodeOut(scope, gradients)) + .Input(ops::AsNodeOut(scope, features)); + scope.UpdateBuilder(&builder); + scope.UpdateStatus(builder.Finalize(scope.graph(), &relu_grad)); + TF_ASSERT_OK(scope.status()); + TF_ASSERT_OK(scope.DoShapeInference(relu_grad)); + auto retval = ops::_Retval(scope.WithOpName("retval"), + Output(relu_grad, 0), 0); + std::unique_ptr graph(new Graph(OpRegistry::Global())); + TF_ASSERT_OK(scope.ToGraph(graph.get())); + + std::vector args(2); + for (XlaCompiler::Argument& arg : args) { + arg.kind = XlaCompiler::Argument::kParameter; + arg.type = DT_FLOAT; + arg.shape = xla::ShapeUtil::MakeShape( + xla::F32, {8, 4}, + std::vector{xla::DExpr::Var(1), xla::DExpr::Const(4)}); + } + + XlaCompiler compiler(DefaultOptions()); + XlaCompiler::CompilationResult result; + TF_ASSERT_OK(compiler.CompileGraph(XlaCompiler::CompileOptions(), + "mlir_relu_grad", std::move(graph), args, + &result)); + + const xla::Shape& result_shape = + xla::ShapeUtil::GetSubshape(result.xla_output_shape, {0}); + EXPECT_FALSE(result_shape.is_dynamic_dimension(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::Const(4))); + + TF_ASSERT_OK_AND_ASSIGN(auto module, + LoadModuleFromHloProto(result.computation->proto())); + for (const xla::HloComputation* computation : module->computations()) { + for (const xla::HloInstruction* instruction : + computation->instructions()) { + EXPECT_NE(instruction->opcode(), xla::HloOpcode::kSetDimensionSize); + } + } +} + +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 = xla::ShapeUtil::MakeShape( + xla::S32, {4}, std::vector{xla::DExpr::Var(1)}); + + 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 +3620,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 +3670,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/compiler/tf2xla/xla_op_registry.cc b/tensorflow/compiler/tf2xla/xla_op_registry.cc index 445065971f2a6a..92493928bc2d6f 100644 --- a/tensorflow/compiler/tf2xla/xla_op_registry.cc +++ b/tensorflow/compiler/tf2xla/xla_op_registry.cc @@ -415,6 +415,18 @@ std::vector XlaOpRegistry::DeviceKernels( return ops; } +/* static */ bool XlaOpRegistry::IsMlirXlaOp(absl::string_view op) { + XlaOpRegistry& registry = Instance(); + mutex_lock lock(registry.mutex_); + auto it = registry.ops_.find(string(op)); + if (it == registry.ops_.end() || it->second.empty()) { + return false; + } + return absl::c_all_of(it->second, [](const auto& registration) { + return registration->uses_mlir_kernel; + }); +} + /*static*/ const std::unordered_set* XlaOpRegistry::CompileTimeConstantInputArgNames(const string& op) { XlaOpRegistry& registry = Instance(); @@ -604,8 +616,9 @@ XlaOpRegistrationBuilder& XlaOpRegistrationBuilder::Label(std::string label) { } std::unique_ptr XlaOpRegistrationBuilder::Build( - XlaOpRegistry::Factory factory) { + XlaOpRegistry::Factory factory, bool uses_mlir_kernel) { registration_->factory = factory; + registration_->uses_mlir_kernel = uses_mlir_kernel; return std::move(registration_); } diff --git a/tensorflow/compiler/tf2xla/xla_op_registry.h b/tensorflow/compiler/tf2xla/xla_op_registry.h index 5eaf0fb2d42bfa..b6abc2a39f71f2 100644 --- a/tensorflow/compiler/tf2xla/xla_op_registry.h +++ b/tensorflow/compiler/tf2xla/xla_op_registry.h @@ -19,6 +19,7 @@ limitations under the License. #include #include #include +#include #include #include @@ -44,6 +45,8 @@ limitations under the License. namespace tensorflow { +class MlirXlaOpKernel; + // Names of the XLA compilation devices. These are not user-visible, and are // used internally by the Tensorflow/XLA bridge to perform symbolic execution of // a Tensorflow graph. @@ -233,6 +236,10 @@ class XlaOpRegistry { // Returns all operations for which there are XLA kernels on any device. static std::vector GetAllRegisteredOps(); + // Returns true if the operation is lowered exclusively through + // MlirXlaOpKernel. + static bool IsMlirXlaOp(absl::string_view op); + // Returns (via `result`) the indices of inputs to `node_def` that must be // compile-time constants. Returns an empty vector if the op is not // registered. @@ -339,6 +346,8 @@ class XlaOpRegistry { // operands and not their values. bool is_metadata_op = false; + bool uses_mlir_kernel = false; + std::string label; // Factory used to build OpKernels that perform symbolic execution. @@ -383,6 +392,9 @@ class XlaOpRegistry { #define REGISTER_XLA_OP(NAME, OP) \ REGISTER_XLA_OP_UNIQ_HELPER(__COUNTER__, NAME, OP) +#define REGISTER_XLA_OP_FACTORY(NAME, FACTORY) \ + REGISTER_XLA_OP_FACTORY_UNIQ_HELPER(__COUNTER__, NAME, FACTORY) + #define REGISTER_XLA_CONV_OP(BUILDER, OP) \ REGISTER_XLA_OP(BUILDER.TypeConstraint("T", GetXlaConvTypesForNonGpu()), OP) \ REGISTER_XLA_OP(BUILDER.TypeConstraint("T", GetXlaConvTypesForGpu()) \ @@ -431,7 +443,7 @@ class XlaOpRegistrationBuilder { XlaOpRegistrationBuilder& Label(std::string label); std::unique_ptr Build( - XlaOpRegistry::Factory factory); + XlaOpRegistry::Factory factory, bool uses_mlir_kernel = false); private: XlaOpRegistrationBuilder(absl::string_view name); @@ -454,11 +466,19 @@ class XlaOpRegistrar { #define REGISTER_XLA_OP_UNIQ_HELPER(COUNTER, BUILDER, OP) \ REGISTER_XLA_OP_UNIQ(COUNTER, BUILDER, OP) +#define REGISTER_XLA_OP_FACTORY_UNIQ_HELPER(COUNTER, BUILDER, FACTORY) \ + REGISTER_XLA_OP_FACTORY_UNIQ(COUNTER, BUILDER, FACTORY) + #define REGISTER_XLA_OP_UNIQ(CTR, BUILDER, OP) \ static ::tensorflow::XlaOpRegistrar xla_op_registrar__body__##CTR##__object( \ ::tensorflow::XlaOpRegistrationBuilder::BUILDER.Build( \ [](::tensorflow::OpKernelConstruction* context) \ - -> ::tensorflow::OpKernel* { return new OP(context); })); + -> ::tensorflow::OpKernel* { return new OP(context); }, \ + std::is_same_v)); + +#define REGISTER_XLA_OP_FACTORY_UNIQ(CTR, BUILDER, FACTORY) \ + static ::tensorflow::XlaOpRegistrar xla_op_registrar__body__##CTR##__object( \ + ::tensorflow::XlaOpRegistrationBuilder::BUILDER.Build(FACTORY)); class XlaBackendRegistrar { public: diff --git a/tensorflow/core/common_runtime/BUILD b/tensorflow/core/common_runtime/BUILD index 301015eba61fe8..f99259ad3e3862 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,13 @@ 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:pack_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.cc b/tensorflow/core/common_runtime/constant_folding.cc index e4427877c33cce..48a36be781fd18 100644 --- a/tensorflow/core/common_runtime/constant_folding.cc +++ b/tensorflow/core/common_runtime/constant_folding.cc @@ -53,6 +53,8 @@ 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 +82,430 @@ bool GetShapeFromDirectDynamicSource(const Node* node, GetShapeFromArgNode(node, out_shape); } +bool TryGetContentsProtoAttr(const AttrSlice& attrs, + TensorShapeProto* out_contents) { + string serialized_contents; + if (!GetNodeAttr(attrs, kUserInferredValueContentsAttrName, + &serialized_contents) + .ok()) { + return false; + } + out_contents->Clear(); + return out_contents->ParseFromString(serialized_contents); +} + +bool HasTransitiveDynamicShapeContents( + 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) && + 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, visiting)) { + return (*memo)[node] = true; + } + } + + return (*memo)[node] = false; +} + +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)); + ExpressionProto* output_expression = output_contents->add_expressions(); + if (index < input_contents.expressions_size()) { + output_expression->CopyFrom(input_contents.expressions(index)); + } else { + output_expression->set_constant_value(input_contents.dim(index).size()); + } +} + +void AppendScalarConstantContent(int64_t value, + TensorShapeProto* output_contents) { + output_contents->add_dim()->set_size(value); + output_contents->add_expressions()->set_constant_value(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, + absl::flat_hash_set* visiting) { + out_contents->Clear(); + if (output_index != 0) { + return false; + } + if (!visiting->insert(node).second) { + return false; + } + + if (TryGetContentsProtoAttr(node->attrs(), out_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, visiting); + }; + + 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; +} + +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 = @@ -272,17 +698,21 @@ bool IsConstantFoldable( shape_map, const std::function& consider, int64_t max_constant_size_in_bytes, - std::unordered_map>* shape_replacement_map) { + std::unordered_map>* shape_replacement_map, + std::unordered_map* dynamic_contents_memo) { + TensorShapeProto exact_contents; + const bool has_exact_contents = + TryGetFoldedValueContents(n, 0, &exact_contents); 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(); + const bool has_dynamic = GetShapeFromDirectDynamicSource(n, &dynamic_shape); + const bool is_shape_derived = + n->attrs().FindByString(kXlaShapeDerivedAttrName) != nullptr; + absl::flat_hash_set dynamic_contents_visiting; + const bool has_transitive_dynamic_contents = + 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; } if (n->IsConstant()) { @@ -367,10 +797,11 @@ void ConsiderConstantFoldableNode( Node* n, const ConstantFoldingOptions& opts, std::vector* nodes, std::unordered_map>* constant_control_deps, std::unordered_map>* shape_replacement_map, + std::unordered_map* dynamic_contents_memo, bool* internal_node_inserted) { if (!IsConstantFoldable(n, opts.shape_map, opts.consider, opts.max_constant_size_in_bytes, - shape_replacement_map)) { + shape_replacement_map, dynamic_contents_memo)) { return; } // A node is constant provided all of its non-control incoming Tensors come @@ -426,13 +857,15 @@ void FindConstantFoldableNodes( std::unordered_map>* constant_control_deps, std::unordered_map>* shape_replacement_map) { bool internal_node_inserted = false; + std::unordered_map dynamic_contents_memo; // Walk the nodes in data flow order. ReverseDFS( *graph, nullptr, [nodes, constant_control_deps, shape_replacement_map, - &internal_node_inserted, &opts](Node* n) { + &dynamic_contents_memo, &internal_node_inserted, &opts](Node* n) { ConsiderConstantFoldableNode(n, opts, nodes, constant_control_deps, shape_replacement_map, + &dynamic_contents_memo, &internal_node_inserted); }, NodeComparatorName()); @@ -492,6 +925,9 @@ 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 +941,10 @@ 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.SerializeAsString()); + } NodeDef def; CHECK(builder.Finalize(&def).ok()); Node* constant_node; @@ -595,6 +1035,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, @@ -626,6 +1072,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 +1098,11 @@ 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.SerializeAsString()); + } if (partition_device) { builder.Device(partition_device->name()); } diff --git a/tensorflow/core/common_runtime/constant_folding_test.cc b/tensorflow/core/common_runtime/constant_folding_test.cc index 481a85add4893c..0702649e5640fa 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,332 @@ 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_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_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_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_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, 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); + 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)); + + 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(); + ASSERT_TRUE(folded->IsConstant()); + 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, FoldPackKeepsSymbolicContentsAligned) { + 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 dynamic_value = ops::GatherV2( + s.WithOpName("dynamic_value"), shape, + ops::Const(s.WithOpName("index"), 0), + ops::Const(s.WithOpName("axis"), 0)); + auto static_value = ops::Const(s.WithOpName("static_value"), 16); + auto pack = ops::Stack(s.WithOpName("pack"), + OutputList{static_value, dynamic_value}, + ops::Stack::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(); + const Edge* send_input = nullptr; + TF_ASSERT_OK(index_by_name.at("send")->input_edge(0, &send_input)); + Node* folded = send_input->src(); + ASSERT_TRUE(folded->IsConstant()); + ExpectNodeEqual(folded, {16, 977}, {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_FALSE(IsDynamicDimExpr(contents_proto.expressions(0))); + EXPECT_TRUE(IsDynamicDimExpr(contents_proto.expressions(1))); + EXPECT_EQ(contents_proto.dim(0).size(), 16); + EXPECT_EQ(contents_proto.dim(1).size(), 977); +} + +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_FALSE(send_input->src()->IsConstant()); +} + TEST_F(ConstantFoldingTest, NoReplacePartialOutput) { Graph g(OpRegistry::Global()); { diff --git a/tensorflow/core/framework/BUILD b/tensorflow/core/framework/BUILD index 57b0b0c39ecb40..79f4bd855fd10b 100644 --- a/tensorflow/core/framework/BUILD +++ b/tensorflow/core/framework/BUILD @@ -748,6 +748,18 @@ cc_library( alwayslink = 1, ) +tf_cc_test( + name = "tensor_shape_expr_test", + size = "small", + srcs = ["tensor_shape_expr_test.cc"], + deps = [ + ":tensor_shape", + ":tensor_shape_expr", + "//tensorflow/core/platform:test", + "//tensorflow/core:test_main", + ], +) + cc_library( name = "resource_base", hdrs = ["resource_base.h"], diff --git a/tensorflow/core/framework/shape_inference.cc b/tensorflow/core/framework/shape_inference.cc index b1336bf4398844..de3e8b5104a05e 100644 --- a/tensorflow/core/framework/shape_inference.cc +++ b/tensorflow/core/framework/shape_inference.cc @@ -248,10 +248,16 @@ void InferenceContext::ShapeHandleToProto(ShapeHandle handle, dim_shape->set_size(Value(dim)); } else { dim_shape->set_size(-1); - // Serialize expression if available. - if (DimExpr* expr = GetDimExpr(dim)) { - expr->ToProto(dim_shape->mutable_expr()); - } + } + + ExpressionProto* expression = proto->add_expressions(); + if (DimExpr* expr = GetDimExpr(dim)) { + DimExprToProto(*expr, expression); + } else if (ValueKnown(dim)) { + expression->set_constant_value(Value(dim)); + } else { + DimExprToProto(DimExpr::Unknown(xla::kMissingExpressionSentinel), + expression); } } } @@ -298,7 +304,8 @@ DimExpr* InferenceContext::GetDimExpr(DimensionHandle d) const { } DimExpr* InferenceContext::MakeConstExpr(int64_t v) { - return shape_manager_.OwnExpr(std::make_unique(v)); + return shape_manager_.OwnExpr( + std::make_unique(DimExpr::Const(v))); } DimExpr* InferenceContext::ExprForDim(DimensionHandle d) { @@ -976,13 +983,15 @@ absl::Status InferenceContext::MakeShapeFromShapeProto( // Known dimension dims.push_back(MakeDim(dim_proto.size())); } else { - // Unknown dimension - check for expression - if (dim_proto.has_expr() && dim_proto.expr().node_type_case() != - ExpressionProto::NODE_TYPE_NOT_SET) { + // Unknown dimension - check for expression. + if (i < proto.expressions_size() && + proto.expressions(i).node_type_case() != + ExpressionProto::NODE_TYPE_NOT_SET) { // Deserialize expression - std::unique_ptr expr = DimExpr::FromProto(dim_proto.expr()); + DimExpr expr = DimExprFromProto(proto.expressions(i)); if (expr) { - DimExpr* owned = shape_manager_.OwnExpr(std::move(expr)); + DimExpr* owned = shape_manager_.OwnExpr( + std::make_unique(std::move(expr))); dims.push_back(shape_manager_.MakeDim(kUnknownDim,/*dynamic_ratio */ 0, owned)); } else { dims.push_back(UnknownDim()); @@ -1128,8 +1137,8 @@ absl::Status InferenceContext::Divide(DimensionHandle dividend, DimExpr* rhs = divisor.dim.IsSet() ? ExprForDim(divisor.dim) : MakeConstExpr(divisor.val); if (lhs && rhs) { - DimExpr* node = shape_manager_.OwnExpr( - std::make_unique(lhs, rhs)); + DimExpr* node = + shape_manager_.OwnExpr(std::make_unique(*lhs / *rhs)); *out = shape_manager_.MakeDim(kUnknownDim, /*dynamic_ratio*/0, node); } else { *out = UnknownDim(); // Can't form expr. @@ -1171,7 +1180,8 @@ absl::Status InferenceContext::Add(DimensionHandle first, second.dim.IsSet() ? ExprForDim(second.dim) : MakeConstExpr(second.val); if (lhs && rhs) { - DimExpr* node = shape_manager_.OwnExpr(std::make_unique(lhs, rhs)); + DimExpr* node = + shape_manager_.OwnExpr(std::make_unique(*lhs + *rhs)); *out = shape_manager_.MakeDim(kUnknownDim, /*dynamic_ratio*/ 0, node); } else { *out = UnknownDim(); // Can't form expr. @@ -1206,7 +1216,8 @@ absl::Status InferenceContext::Subtract(DimensionHandle first, DimExpr* rhs = second.dim.IsSet() ? ExprForDim(second.dim) : MakeConstExpr(second.val); if (lhs && rhs) { - DimExpr* node = shape_manager_.OwnExpr(std::make_unique(lhs, rhs)); + DimExpr* node = + shape_manager_.OwnExpr(std::make_unique(*lhs - *rhs)); *out = shape_manager_.MakeDim(kUnknownDim, /*dynamic_ratio*/ 0, node); } else { *out = UnknownDim(); // Can't form expr. @@ -1258,7 +1269,8 @@ absl::Status InferenceContext::Multiply(DimensionHandle first, second.dim.IsSet() ? ExprForDim(second.dim) : MakeConstExpr(second.val); if (lhs && rhs) { - DimExpr* node = shape_manager_.OwnExpr(std::make_unique(lhs, rhs)); + DimExpr* node = + shape_manager_.OwnExpr(std::make_unique(*lhs * *rhs)); *out = shape_manager_.MakeDim(kUnknownDim, /*dynamic_ratio*/ 0, node); } else { *out = UnknownDim(); // Can't form expr. diff --git a/tensorflow/core/framework/shape_inference_test.cc b/tensorflow/core/framework/shape_inference_test.cc index c5dbc299b86540..9a10a05ed65163 100644 --- a/tensorflow/core/framework/shape_inference_test.cc +++ b/tensorflow/core/framework/shape_inference_test.cc @@ -1035,6 +1035,35 @@ TEST_F(ShapeInferenceTest, KnownShapeToProto) { EXPECT_FALSE(proto.unknown_rank()); EXPECT_EQ(3, proto.dim_size()); EXPECT_EQ(1, proto.dim(0).size()); + ASSERT_EQ(3, proto.expressions_size()); + EXPECT_EQ(1, proto.expressions(0).constant_value()); + EXPECT_EQ(2, proto.expressions(1).constant_value()); + EXPECT_EQ(3, proto.expressions(2).constant_value()); + EXPECT_FALSE(proto.dim(0).has_expr()); +} + +TEST_F(ShapeInferenceTest, DynamicExpressionShapeProtoRoundTrip) { + NodeDef def; + std::vector empty; + InferenceContext c(kVersion, def, MakeOpDef(0, 2), empty, {}, {}, {}); + + TensorShapeProto input_proto; + input_proto.add_dim()->set_size(-1); + DimExprToProto(DimExpr::Var(7), input_proto.add_expressions()); + + ShapeHandle shape; + TF_ASSERT_OK(c.MakeShapeFromShapeProto(input_proto, &shape)); + TensorShapeProto proto; + c.ShapeHandleToProto(shape, &proto); + + ASSERT_EQ(1, proto.expressions_size()); + EXPECT_EQ(DimExpr::Var(7), DimExprFromProto(proto.expressions(0))); + EXPECT_FALSE(proto.dim(0).has_expr()); + + ShapeHandle restored; + TF_ASSERT_OK(c.MakeShapeFromShapeProto(proto, &restored)); + ASSERT_NE(nullptr, c.GetDimExpr(c.Dim(restored, 0))); + EXPECT_EQ(DimExpr::Var(7), *c.GetDimExpr(c.Dim(restored, 0))); } TEST_F(ShapeInferenceTest, UnknownShapeToProto) { diff --git a/tensorflow/core/framework/tensor_shape.cc b/tensorflow/core/framework/tensor_shape.cc index fd2606224b3429..4b06e0861d86cf 100644 --- a/tensorflow/core/framework/tensor_shape.cc +++ b/tensorflow/core/framework/tensor_shape.cc @@ -28,91 +28,6 @@ limitations under the License. namespace tensorflow { -namespace { - -const bool kTensorShapeExpressionsEnabled = TensorShapeExpressionsEnabled(); - -xla::DExpr DExprFromProto(const ExpressionProto& proto) { - switch (proto.node_type_case()) { - case ExpressionProto::kConstantValue: - return xla::DExpr::Const(proto.constant_value()); - case ExpressionProto::kVariableId: - return xla::DExpr::Var(proto.variable_id()); - case ExpressionProto::kAddNode: { - const auto& add = proto.add_node(); - return DExprFromProto(add.lhs()) + DExprFromProto(add.rhs()); - } - case ExpressionProto::kSubNode: { - const auto& sub = proto.sub_node(); - return DExprFromProto(sub.lhs()) - DExprFromProto(sub.rhs()); - } - case ExpressionProto::kMulNode: { - const auto& mul = proto.mul_node(); - return DExprFromProto(mul.lhs()) * DExprFromProto(mul.rhs()); - } - case ExpressionProto::kDivNode: { - const auto& div = proto.div_node(); - return DExprFromProto(div.lhs()) / DExprFromProto(div.rhs()); - } - case ExpressionProto::NODE_TYPE_NOT_SET: - default: - return xla::DExpr::Unknown(xla::kMissingExpressionSentinel); - } -} - -void ExprToProto(const xla::DExpr& expr, ExpressionProto* proto) { - if (!expr) return; - switch (expr.kind()) { - case xla::DExpr::Kind::kUnknown: - return; - case xla::DExpr::Kind::kConstant: - proto->set_constant_value(expr->get_val()); - return; - case xla::DExpr::Kind::kVariable: - proto->set_variable_id( - static_cast(*expr.get()).get_id()); - return; - case xla::DExpr::Kind::kAdd: { - auto* add = proto->mutable_add_node(); - const auto& node = static_cast(*expr.get()); - ExprToProto(xla::DExpr(node.get_lhs()->clone()), - add->mutable_lhs()); - ExprToProto(xla::DExpr(node.get_rhs()->clone()), - add->mutable_rhs()); - return; - } - case xla::DExpr::Kind::kSub: { - auto* sub = proto->mutable_sub_node(); - const auto& node = static_cast(*expr.get()); - ExprToProto(xla::DExpr(node.get_lhs()->clone()), - sub->mutable_lhs()); - ExprToProto(xla::DExpr(node.get_rhs()->clone()), - sub->mutable_rhs()); - return; - } - case xla::DExpr::Kind::kMul: { - auto* mul = proto->mutable_mul_node(); - const auto& node = static_cast(*expr.get()); - ExprToProto(xla::DExpr(node.get_lhs()->clone()), - mul->mutable_lhs()); - ExprToProto(xla::DExpr(node.get_rhs()->clone()), - mul->mutable_rhs()); - return; - } - case xla::DExpr::Kind::kDiv: { - auto* div = proto->mutable_div_node(); - const auto& node = static_cast(*expr.get()); - ExprToProto(xla::DExpr(node.get_lhs()->clone()), - div->mutable_lhs()); - ExprToProto(xla::DExpr(node.get_rhs()->clone()), - div->mutable_rhs()); - return; - } - } -} - -} // namespace - std::string ExprToString(const xla::DExpr& e) { if (!e && !e.is_unknown()) return ""; xla::StringPrinter printer; @@ -246,9 +161,9 @@ 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)); + AddExpression(DimExprFromProto(e)); } } } @@ -290,9 +205,9 @@ absl::Status TensorShapeBase::BuildTensorShapeBase( } } } - if (kTensorShapeExpressionsEnabled) { + if (TensorShapeExpressionsEnabled()) { for (const auto& e : proto.expressions()) { - out->AddExpression(DExprFromProto(e)); + out->AddExpression(DimExprFromProto(e)); } } } @@ -480,7 +395,7 @@ void TensorShapeRep::Clear() { } void TensorShapeRep::set_expression(int d, xla::DExpr expr) { - if (!kTensorShapeExpressionsEnabled) { + if (!TensorShapeExpressionsEnabled()) { expressions_.clear(); return; } @@ -493,7 +408,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,10 +418,11 @@ void TensorShapeRep::AddExpression(xla::DExpr expr) { } void TensorShapeRep::set_expressions(std::vector exprs) { - if (!kTensorShapeExpressionsEnabled) { + if (!TensorShapeExpressionsEnabled()) { expressions_.clear(); return; } + CHECK_LE(exprs.size(), ndims_byte()); for (auto& expr : exprs) { if (!expr) expr = xla::DExpr::Unknown(xla::kMissingExpressionSentinel); } @@ -838,10 +754,10 @@ void TensorShapeBase::RemoveDimRange(int begin, int end) { } ClearAllButDataType(); - set_expressions(new_exprs); for (auto dval : vals) { AddDim(dval); } + set_expressions(new_exprs); TF_CHECK_OK(RecomputeNumElements()); } @@ -897,7 +813,6 @@ absl::Status TensorShapeBase::RemoveDimRangeWithStatus(int begin, } ClearAllButDataType(); - set_expressions(new_exprs); absl::Status s = absl::OkStatus(); for (auto dval : vals) { @@ -906,6 +821,7 @@ absl::Status TensorShapeBase::RemoveDimRangeWithStatus(int begin, return s; } } + set_expressions(new_exprs); return RecomputeNumElements(); } @@ -926,10 +842,10 @@ 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); + DimExprToProto(get_expression(i), eproto); } } } @@ -993,13 +909,12 @@ 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()) { if (!first) strings::StrAppend(&s, ","); - auto exp = DExprFromProto(e); - strings::StrAppend(&s, ExprToString(exp)); + strings::StrAppend(&s, ExprToString(DimExprFromProto(e))); first = false; } strings::StrAppend(&s, ">"); 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 baca4df733d63f..5acb14ef9aa880 100644 --- a/tensorflow/core/framework/tensor_shape_expr.cc +++ b/tensorflow/core/framework/tensor_shape_expr.cc @@ -1,8 +1,11 @@ #include "tensorflow/core/framework/tensor_shape_expr.h" +#include +#include #include #include "xla/parse_flags_from_env.h" +#include "xla/printer.h" #include "xla/tsl/util/command_line_flags.h" namespace tensorflow { @@ -19,233 +22,178 @@ bool ParseTensorShapeExpressionsEnabled() { return tf_xla_enable_dynamic_sizes; } -} // namespace - -bool TensorShapeExpressionsEnabled() { - static const bool enabled = ParseTensorShapeExpressionsEnabled(); - return enabled; -} - -bool IsDynamicDimExpr(const ExpressionProto& proto) { - switch (proto.node_type_case()) { - case ExpressionProto::kVariableId: - return true; - case ExpressionProto::kAddNode: - return IsDynamicDimExpr(proto.add_node().lhs()) || - IsDynamicDimExpr(proto.add_node().rhs()); - case ExpressionProto::kSubNode: - return IsDynamicDimExpr(proto.sub_node().lhs()) || - IsDynamicDimExpr(proto.sub_node().rhs()); - case ExpressionProto::kMulNode: - return IsDynamicDimExpr(proto.mul_node().lhs()) || - IsDynamicDimExpr(proto.mul_node().rhs()); - case ExpressionProto::kDivNode: - return IsDynamicDimExpr(proto.div_node().lhs()) || - IsDynamicDimExpr(proto.div_node().rhs()); - case ExpressionProto::kConstantValue: - case ExpressionProto::NODE_TYPE_NOT_SET: - return false; - } +std::optional& TensorShapeExpressionsEnabledOverride() { + static auto* enabled_override = new std::optional(); + return *enabled_override; } -bool HasDynamicDimExprs(const TensorShapeProto& proto) { - for (const auto& expr : proto.expressions()) { - if (IsDynamicDimExpr(expr)) { - return true; +void DynExprToTensorFlowProto(const xla::DynExpr& expr, + ExpressionProto* proto) { + proto->Clear(); + switch (expr.kind()) { + case xla::DExpr::Kind::kUnknown: + return; + case xla::DExpr::Kind::kConstant: + proto->set_constant_value( + static_cast(expr).get_val()); + return; + case xla::DExpr::Kind::kVariable: + proto->set_variable_id( + static_cast(expr).get_id()); + return; + case xla::DExpr::Kind::kAdd: { + const auto& add = static_cast(expr); + DynExprToTensorFlowProto(*add.get_lhs(), + proto->mutable_add_node()->mutable_lhs()); + DynExprToTensorFlowProto(*add.get_rhs(), + proto->mutable_add_node()->mutable_rhs()); + return; } - } - return false; -} - -std::unique_ptr DimExpr::Cons(int64_t val) { - return std::make_unique(val); -} - -std::unique_ptr DimExpr::Var(int32_t id) { - return std::make_unique(id); -} - -std::string DimExpr::DebugString() const { - ExpressionProto proto; - ToProto(&proto); - return proto.DebugString(); -} - -static bool EqualsImpl(const DimExpr* a, const DimExpr* b) { - if (a == b) return true; - if (a == nullptr || b == nullptr) return false; - if (a->kind() != b->kind()) return false; - - switch (a->kind()) { - case DimExpr::Kind::kConstant: { - auto* ac = static_cast(a); - auto* bc = static_cast(b); - return ac->value() == bc->value(); + case xla::DExpr::Kind::kSub: { + const auto& sub = static_cast(expr); + DynExprToTensorFlowProto(*sub.get_lhs(), + proto->mutable_sub_node()->mutable_lhs()); + DynExprToTensorFlowProto(*sub.get_rhs(), + proto->mutable_sub_node()->mutable_rhs()); + return; } - case DimExpr::Kind::kVariable: { - auto* av = static_cast(a); - auto* bv = static_cast(b); - return av->id() == bv->id(); + case xla::DExpr::Kind::kMul: { + const auto& mul = static_cast(expr); + DynExprToTensorFlowProto(*mul.get_lhs(), + proto->mutable_mul_node()->mutable_lhs()); + DynExprToTensorFlowProto(*mul.get_rhs(), + proto->mutable_mul_node()->mutable_rhs()); + return; } - case DimExpr::Kind::kAdd: { - auto* aa = static_cast(a); - auto* ba = static_cast(b); - return EqualsImpl(aa->lhs(), ba->lhs()) && - EqualsImpl(aa->rhs(), ba->rhs()); + case xla::DExpr::Kind::kDiv: { + const auto& div = static_cast(expr); + DynExprToTensorFlowProto(*div.get_lhs(), + proto->mutable_div_node()->mutable_lhs()); + DynExprToTensorFlowProto(*div.get_rhs(), + proto->mutable_div_node()->mutable_rhs()); + return; } - case DimExpr::Kind::kSub: { - auto* as = static_cast(a); - auto* bs = static_cast(b); - return EqualsImpl(as->lhs(), bs->lhs()) && - EqualsImpl(as->rhs(), bs->rhs()); + case xla::DExpr::Kind::kMax: { + const auto& max = static_cast(expr); + DynExprToTensorFlowProto(*max.get_lhs(), + proto->mutable_max_node()->mutable_lhs()); + DynExprToTensorFlowProto(*max.get_rhs(), + proto->mutable_max_node()->mutable_rhs()); + return; } - case DimExpr::Kind::kMul: { - auto* am = static_cast(a); - auto* bm = static_cast(b); - return EqualsImpl(am->lhs(), bm->lhs()) && - EqualsImpl(am->rhs(), bm->rhs()); + case xla::DExpr::Kind::kGt: { + const auto& gt = static_cast(expr); + DynExprToTensorFlowProto(*gt.get_lhs(), + proto->mutable_gt_node()->mutable_lhs()); + DynExprToTensorFlowProto(*gt.get_rhs(), + proto->mutable_gt_node()->mutable_rhs()); + return; } - case DimExpr::Kind::kDiv: { - auto* ad = static_cast(a); - auto* bd = static_cast(b); - return EqualsImpl(ad->lhs(), bd->lhs()) && - EqualsImpl(ad->rhs(), bd->rhs()); + case xla::DExpr::Kind::kSelect: { + const auto& select = static_cast(expr); + DynExprToTensorFlowProto( + *select.get_pred(), proto->mutable_select_node()->mutable_pred()); + DynExprToTensorFlowProto( + *select.get_on_true(), + proto->mutable_select_node()->mutable_on_true()); + DynExprToTensorFlowProto( + *select.get_on_false(), + proto->mutable_select_node()->mutable_on_false()); + return; } } +} - return false; +} // namespace + +bool TensorShapeExpressionsEnabled() { + if (TensorShapeExpressionsEnabledOverride().has_value()) { + return *TensorShapeExpressionsEnabledOverride(); + } + static const bool enabled = ParseTensorShapeExpressionsEnabled(); + return enabled; } -bool DimExpr::Equals(const DimExpr* a, const DimExpr* b) { - return EqualsImpl(a, b); +void SetTensorShapeExpressionsEnabledForTesting(std::optional enabled) { + TensorShapeExpressionsEnabledOverride() = enabled; } -std::unique_ptr DimExpr::FromProto(const ExpressionProto& proto) { +DimExpr DimExprFromProto(const ExpressionProto& proto) { switch (proto.node_type_case()) { case ExpressionProto::kConstantValue: - return DimExpr::Cons(proto.constant_value()); + return DimExpr::Const(proto.constant_value()); case ExpressionProto::kVariableId: return DimExpr::Var(proto.variable_id()); case ExpressionProto::kAddNode: { - auto lhs = FromProto(proto.add_node().lhs()); - auto rhs = FromProto(proto.add_node().rhs()); - // Note: These are owning pointers, but ExprAdd takes raw pointers. - // The caller must manage lifetime appropriately. - return std::make_unique(lhs.release(), rhs.release()); + return DimExprFromProto(proto.add_node().lhs()) + + DimExprFromProto(proto.add_node().rhs()); } case ExpressionProto::kSubNode: { - auto lhs = FromProto(proto.sub_node().lhs()); - auto rhs = FromProto(proto.sub_node().rhs()); - return std::make_unique(lhs.release(), rhs.release()); + return DimExprFromProto(proto.sub_node().lhs()) - + DimExprFromProto(proto.sub_node().rhs()); } case ExpressionProto::kMulNode: { - auto lhs = FromProto(proto.mul_node().lhs()); - auto rhs = FromProto(proto.mul_node().rhs()); - return std::make_unique(lhs.release(), rhs.release()); + return DimExprFromProto(proto.mul_node().lhs()) * + DimExprFromProto(proto.mul_node().rhs()); } case ExpressionProto::kDivNode: { - auto lhs = FromProto(proto.div_node().lhs()); - auto rhs = FromProto(proto.div_node().rhs()); - return std::make_unique(lhs.release(), rhs.release()); + return DimExprFromProto(proto.div_node().lhs()) / + DimExprFromProto(proto.div_node().rhs()); + } + case ExpressionProto::kMaxNode: { + return DimExpr::Max(DimExprFromProto(proto.max_node().lhs()), + DimExprFromProto(proto.max_node().rhs())); + } + case ExpressionProto::kGtNode: { + return DimExpr::Gt(DimExprFromProto(proto.gt_node().lhs()), + DimExprFromProto(proto.gt_node().rhs())); + } + case ExpressionProto::kSelectNode: { + return DimExpr::Select( + DimExprFromProto(proto.select_node().pred()), + DimExprFromProto(proto.select_node().on_true()), + DimExprFromProto(proto.select_node().on_false())); } case ExpressionProto::NODE_TYPE_NOT_SET: default: - return nullptr; + return DimExpr::Unknown(xla::kMissingExpressionSentinel); } } -DimExpr* SimplifyExpr(DimExpr* expr, - std::vector>* arena) { - if (!expr) return nullptr; - - auto own = [arena](std::unique_ptr e) -> DimExpr* { - DimExpr* ptr = e.get(); - arena->push_back(std::move(e)); - return ptr; - }; - - switch (expr->kind()) { - case DimExpr::Kind::kConstant: - case DimExpr::Kind::kVariable: - return expr; - - case DimExpr::Kind::kAdd: { - auto* add = static_cast(expr); - DimExpr* lhs = SimplifyExpr(add->lhs(), arena); - DimExpr* rhs = SimplifyExpr(add->rhs(), arena); - - // Constant folding - if (lhs->IsConstant() && rhs->IsConstant()) { - return own(DimExpr::Cons(lhs->ConstantValue() + rhs->ConstantValue())); - } - - // x + 0 → x - if (rhs->IsConstant() && rhs->ConstantValue() == 0) return lhs; - if (lhs->IsConstant() && lhs->ConstantValue() == 0) return rhs; - - return own(std::make_unique(lhs, rhs)); - } - - case DimExpr::Kind::kSub: { - auto* sub = static_cast(expr); - DimExpr* lhs = SimplifyExpr(sub->lhs(), arena); - DimExpr* rhs = SimplifyExpr(sub->rhs(), arena); - - // Constant folding - if (lhs->IsConstant() && rhs->IsConstant()) { - return own(DimExpr::Cons(lhs->ConstantValue() - rhs->ConstantValue())); - } - - // x - 0 → x - if (rhs->IsConstant() && rhs->ConstantValue() == 0) return lhs; - - return own(std::make_unique(lhs, rhs)); - } - - case DimExpr::Kind::kMul: { - auto* mul = static_cast(expr); - DimExpr* lhs = SimplifyExpr(mul->lhs(), arena); - DimExpr* rhs = SimplifyExpr(mul->rhs(), arena); - - // Constant folding - if (lhs->IsConstant() && rhs->IsConstant()) { - return own(DimExpr::Cons(lhs->ConstantValue() * rhs->ConstantValue())); - } - - // x * 1 → x - if (rhs->IsConstant() && rhs->ConstantValue() == 1) return lhs; - if (lhs->IsConstant() && lhs->ConstantValue() == 1) return rhs; - - // x * 0 → 0 - if (rhs->IsConstant() && rhs->ConstantValue() == 0) - return own(DimExpr::Cons(0)); - if (lhs->IsConstant() && lhs->ConstantValue() == 0) - return own(DimExpr::Cons(0)); - - return own(std::make_unique(lhs, rhs)); - } +void DimExprToProto(const DimExpr& expr, ExpressionProto* proto) { + if (!expr) { + proto->Clear(); + return; + } + DynExprToTensorFlowProto(*expr, proto); +} - case DimExpr::Kind::kDiv: { - auto* div = static_cast(expr); - DimExpr* lhs = SimplifyExpr(div->lhs(), arena); - DimExpr* rhs = SimplifyExpr(div->rhs(), arena); +std::string DimExprDebugString(const DimExpr& expr) { + if (!expr) return "_"; + xla::StringPrinter printer; + expr->print(&printer); + return std::move(printer).ToString(); +} - // Constant folding (avoid div by zero) - if (lhs->IsConstant() && rhs->IsConstant()) { - int64_t r = rhs->ConstantValue(); - if (r != 0) { - return own(DimExpr::Cons(lhs->ConstantValue() / r)); - } - } +DimExpr* SimplifyExpr(DimExpr* expr, + std::vector>* arena) { + if (expr == nullptr) return nullptr; + auto owned = std::make_unique(expr->simplify()); + DimExpr* result = owned.get(); + arena->push_back(std::move(owned)); + return result; +} - // x / 1 → x - if (rhs->IsConstant() && rhs->ConstantValue() == 1) return lhs; +bool IsDynamicDimExpr(const ExpressionProto& proto) { + DimExpr expr = DimExprFromProto(proto); + return expr && expr->is_dynamic(); +} - return own(std::make_unique(lhs, rhs)); - } +bool HasDynamicDimExprs(const TensorShapeProto& proto) { + for (const auto& expr : proto.expressions()) { + if (IsDynamicDimExpr(expr)) return true; } - - return expr; + return false; } } // namespace tensorflow diff --git a/tensorflow/core/framework/tensor_shape_expr.h b/tensorflow/core/framework/tensor_shape_expr.h index 1c215fda268dcf..4847c8c10b4785 100644 --- a/tensorflow/core/framework/tensor_shape_expr.h +++ b/tensorflow/core/framework/tensor_shape_expr.h @@ -1,223 +1,36 @@ #ifndef TENSORFLOW_CORE_FRAMEWORK_TENSOR_SHAPE_EXPR_H_ #define TENSORFLOW_CORE_FRAMEWORK_TENSOR_SHAPE_EXPR_H_ -#include #include +#include #include -#include +#include #include "tensorflow/core/framework/tensor_shape.pb.h" +#include "xla/shape_expr.h" namespace tensorflow { -// Forward declarations -class Constant; -class Variable; -class ExprAdd; -class ExprSub; -class ExprMul; -class ExprDiv; +// TensorFlow shape inference and XLA use the same owning expression value. +// TensorFlow-specific helpers below only bridge TensorShapeProto's protobuf. +using DimExpr = xla::DExpr; -// DimExpr: Base class for symbolic expressions representing dynamic dimension -// sizes. These expressions form a DAG that tracks how unknown dimensions relate -// to each other through arithmetic operations. -// -// The expression language: -// - 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 -// -// INVARIANT: An unknown dimension is not just -1, it is -1 + Var(sym). -class DimExpr { - public: - enum class Kind : uint8_t { - kConstant, - kVariable, - kAdd, - kSub, - kMul, - kDiv, - }; +DimExpr DimExprFromProto(const ExpressionProto& proto); +void DimExprToProto(const DimExpr& expr, ExpressionProto* proto); +std::string DimExprDebugString(const DimExpr& expr); - virtual ~DimExpr() = default; - - virtual Kind kind() const = 0; - virtual void ToProto(ExpressionProto* proto) const = 0; - - virtual bool IsConstant() const { return false; } - virtual int64_t ConstantValue() const { return 0; } - - // Factory methods - return owning pointers - static std::unique_ptr Cons(int64_t val); - static std::unique_ptr Var(int32_t var_id); - - // Structural equality check - static bool Equals(const DimExpr* a, const DimExpr* b); - - // Build from proto (owns all returned nodes) - static std::unique_ptr FromProto(const ExpressionProto& proto); - - // Debug representation - std::string DebugString() const; - - protected: - DimExpr() = default; -}; - -// Constant expression node: represents a known integer value -class Constant final : public DimExpr { - public: - explicit Constant(int64_t value) : value_(value) {} - - Kind kind() const override { return Kind::kConstant; } - void ToProto(ExpressionProto* proto) const override { - proto->set_constant_value(value_); - } - - bool IsConstant() const override { return true; } - int64_t ConstantValue() const override { return value_; } - - int64_t value() const { return value_; } - - private: - int64_t value_; -}; - -// Variable expression node: represents a symbolic unknown dimension -class Variable final : public DimExpr { - public: - explicit Variable(int32_t id) : id_(id) {} - - Kind kind() const override { return Kind::kVariable; } - void ToProto(ExpressionProto* proto) const override { - proto->set_variable_id(id_); - } - - int32_t id() const { return id_; } - - private: - int32_t id_; -}; - -// Addition expression node -class ExprAdd final : public DimExpr { - public: - ExprAdd(DimExpr* lhs, DimExpr* rhs) : lhs_(lhs), rhs_(rhs) {} - - Kind kind() const override { return Kind::kAdd; } - void ToProto(ExpressionProto* proto) const override { - auto* add_msg = proto->mutable_add_node(); - lhs_->ToProto(add_msg->mutable_lhs()); - rhs_->ToProto(add_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_; -}; - -// Subtraction expression node -class ExprSub final : public DimExpr { - public: - ExprSub(DimExpr* lhs, DimExpr* rhs) : lhs_(lhs), rhs_(rhs) {} - - Kind kind() const override { return Kind::kSub; } - void ToProto(ExpressionProto* proto) const override { - auto* sub_msg = proto->mutable_sub_node(); - lhs_->ToProto(sub_msg->mutable_lhs()); - rhs_->ToProto(sub_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_; -}; - -// Multiplication expression node -class ExprMul final : public DimExpr { - public: - ExprMul(DimExpr* lhs, DimExpr* rhs) : lhs_(lhs), rhs_(rhs) {} - - Kind kind() const override { return Kind::kMul; } - void ToProto(ExpressionProto* proto) const override { - auto* mul_msg = proto->mutable_mul_node(); - lhs_->ToProto(mul_msg->mutable_lhs()); - rhs_->ToProto(mul_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_; -}; - -// Division expression node -class ExprDiv final : public DimExpr { - public: - ExprDiv(DimExpr* lhs, DimExpr* rhs) : lhs_(lhs), rhs_(rhs) {} - - Kind kind() const override { return Kind::kDiv; } - void ToProto(ExpressionProto* proto) const override { - auto* div_msg = proto->mutable_div_node(); - lhs_->ToProto(div_msg->mutable_lhs()); - rhs_->ToProto(div_msg->mutable_rhs()); - } - - bool IsConstant() const override { - return lhs_->IsConstant() && rhs_->IsConstant(); - } - int64_t ConstantValue() const override { - int64_t r = rhs_->ConstantValue(); - return (r == 0) ? 0 : lhs_->ConstantValue() / r; - } - - DimExpr* lhs() const { return lhs_; } - DimExpr* rhs() const { return rhs_; } - - private: - DimExpr* lhs_; - DimExpr* rhs_; -}; - -// 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. +// Simplifies through xla::DExpr and stores the returned value in `arena`. DimExpr* SimplifyExpr(DimExpr* expr, std::vector>* arena); // Returns whether TensorShape should preserve symbolic expressions. The -// Shape-expression support follows the `tf_xla_enable_dynamic_sizes` flag. +// 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/tensorflow/core/framework/tensor_shape_expr_test.cc b/tensorflow/core/framework/tensor_shape_expr_test.cc new file mode 100644 index 00000000000000..55e6fe915332aa --- /dev/null +++ b/tensorflow/core/framework/tensor_shape_expr_test.cc @@ -0,0 +1,90 @@ +#include "tensorflow/core/framework/tensor_shape_expr.h" + +#include +#include +#include + +#include "tensorflow/core/framework/tensor_shape.h" +#include "tensorflow/core/platform/test.h" + +namespace tensorflow { +namespace { + +class TensorShapeExpressionsEnabledTest : public ::testing::Test { + protected: + void SetUp() override { + SetTensorShapeExpressionsEnabledForTesting(true); + } + + void TearDown() override { + SetTensorShapeExpressionsEnabledForTesting(std::nullopt); + } +}; + +TEST(TensorShapeExprTest, UsesXlaCanonicalization) { + ExpressionProto proto; + auto* add = proto.mutable_add_node(); + add->mutable_lhs()->mutable_div_node()->mutable_lhs()->set_variable_id(1); + add->mutable_lhs()->mutable_div_node()->mutable_rhs()->set_constant_value(2); + add->mutable_rhs()->mutable_div_node()->mutable_lhs()->set_variable_id(1); + add->mutable_rhs()->mutable_div_node()->mutable_rhs()->set_constant_value(2); + + auto expr = DimExprFromProto(proto); + EXPECT_EQ(DimExprDebugString(expr.simplify()), "A"); +} + +TEST(TensorShapeExprTest, TensorFlowProtoRoundTripsSharedExpression) { + DimExpr original = (DimExpr::Var(7) + DimExpr::Const(3)) / 2; + ExpressionProto proto; + DimExprToProto(original, &proto); + + auto round_tripped = DimExprFromProto(proto); + EXPECT_TRUE(xla::DynExpr::equal(original.get(), round_tripped.get())); +} + +TEST(TensorShapeExprTest, TensorFlowProtoRoundTripsConditionalExpression) { + DimExpr variable = DimExpr::Var(7); + DimExpr original = DimExpr::Select( + DimExpr::Gt(variable, DimExpr::Const(3)), + DimExpr::Max(variable, DimExpr::Const(8)), + (variable + DimExpr::Const(3)) / 2); + ExpressionProto proto; + DimExprToProto(original, &proto); + + auto round_tripped = DimExprFromProto(proto); + EXPECT_TRUE(xla::DynExpr::equal(original.get(), round_tripped.get())); +} + +TEST(TensorShapeExprTest, SimplifyExprUsesSharedImplementation) { + DimExpr original = + (DimExpr::Var(1) / 2) + (DimExpr::Var(1) / 2); + std::vector> arena; + + DimExpr* simplified = SimplifyExpr(&original, &arena); + + ASSERT_NE(simplified, nullptr); + EXPECT_EQ(DimExprDebugString(*simplified), "A"); +} + +TEST_F(TensorShapeExpressionsEnabledTest, + SetExpressionsRejectsEntriesBeyondRank) { + TensorShape shape({2, 3}); + EXPECT_DEATH( + shape.set_expressions( + {xla::DExpr::Var(1), xla::DExpr::Var(2), xla::DExpr::Var(3)}), + ""); +} + +TEST_F(TensorShapeExpressionsEnabledTest, + RemoveDimRangePreservesRemainingExpressions) { + TensorShape shape({2, 3}); + shape.set_expressions({xla::DExpr::Var(1), xla::DExpr::Var(2)}); + + shape.RemoveDim(0); + + ASSERT_EQ(shape.get_expressions().size(), 1); + EXPECT_TRUE(shape.get_expression(0) == xla::DExpr::Var(2)); +} + +} // namespace +} // namespace tensorflow diff --git a/tensorflow/core/grappler/costs/graph_properties.cc b/tensorflow/core/grappler/costs/graph_properties.cc index 774c628ce07c34..bee993c8a1142c 100644 --- a/tensorflow/core/grappler/costs/graph_properties.cc +++ b/tensorflow/core/grappler/costs/graph_properties.cc @@ -652,11 +652,13 @@ class SymbolicShapeRefiner { explicit SymbolicShapeRefiner( const GraphView& graph, const absl::flat_hash_map>& fed_ports, - const bool aggressive_shape_inference) + const bool aggressive_shape_inference, + const bool enable_dynamic_value_inference) : graph_(graph), function_library_(OpRegistry::Global(), graph.graph()->library()), fed_ports_(fed_ports), - aggressive_shape_inference_(aggressive_shape_inference) { + aggressive_shape_inference_(aggressive_shape_inference), + enable_dynamic_value_inference_(enable_dynamic_value_inference) { graph_def_version_ = graph.graph()->versions().producer(); node_to_context_.reserve(graph.graph()->node_size()); } @@ -945,7 +947,9 @@ class SymbolicShapeRefiner { TF_RETURN_IF_ERROR(gp.InferStatically( /*assume_valid_feeds=*/true, /*aggressive_shape_inference=*/aggressive_shape_inference_, - /*include_tensor_values=*/true)); + /*include_input_tensor_values=*/true, + /*include_output_tensor_values=*/true, + /*enable_dynamic_value_inference=*/enable_dynamic_value_inference_)); // Add return nodes for output shapes. int output = 0; @@ -1420,9 +1424,11 @@ class SymbolicShapeRefiner { if (node->op() == "_Arg") { var_id *= -1; // var_id would be minus when it's argument. - dim = c->UnknownDimWithExpr(DimExpr::Var(var_id)); + dim = c->UnknownDimWithExpr( + std::make_unique(DimExpr::Var(var_id))); } else { - dim = c->UnknownDimWithExpr(DimExpr::Var(var_id)); + dim = c->UnknownDimWithExpr( + std::make_unique(DimExpr::Var(var_id))); } VLOG(1) << "[EXPR] GetUnknownOutputDim: node=" << node->name() << " out=" << index << " dim=" << dim_id << " -> Var(" << var_id @@ -1742,44 +1748,56 @@ class SymbolicShapeRefiner { const bool is_bin = (op == "Sub" || op == "Add" || op == "Mul" || op == "Div"); if (!is_fed) { - if (is_bin) { - if (c->input_tensors_as_shapes_to_propagate.size() < 2) - return absl::OkStatus(); + // Fall through to regular value inference when symbolic shape-value + // propagation does not apply. + if (enable_dynamic_value_inference_ && is_bin && + c->input_tensors_as_shapes_to_propagate.size() >= 2) { auto va = c->input_tensors_as_shapes_to_propagate[0]; auto vb = c->input_tensors_as_shapes_to_propagate[1]; - if (va.SameHandle(tensorflow::shape_inference::ShapeHandle()) || - vb.SameHandle(tensorflow::shape_inference::ShapeHandle())) { - return absl::OkStatus(); - } - - if (!ic->RankKnown(va) || !ic->RankKnown(vb)) return absl::OkStatus(); - if (ic->Rank(va) != ic->Rank(vb)) return absl::OkStatus(); - - std::vector out_elems; - out_elems.reserve(ic->Rank(va)); + if (!va.SameHandle(tensorflow::shape_inference::ShapeHandle()) && + !vb.SameHandle(tensorflow::shape_inference::ShapeHandle()) && + ic->RankKnown(va) && ic->RankKnown(vb) && + ic->Rank(va) == ic->Rank(vb)) { + std::vector out_elems; + out_elems.reserve(ic->Rank(va)); + const auto is_unknown_from_const = [&](DimensionHandle dim) { + return ic->ValueKnown(dim) && + ic->Value(dim) == kUnknownDimFromConst; + }; + bool symbolic_propagation_succeeded = true; - for (int i = 0; i < ic->Rank(va); ++i) { - auto da = ic->Dim(va, i); - auto db = ic->Dim(vb, i); + for (int i = 0; i < ic->Rank(va); ++i) { + auto da = ic->Dim(va, i); + auto db = ic->Dim(vb, i); + if (is_unknown_from_const(da) || is_unknown_from_const(db)) { + symbolic_propagation_succeeded = false; + break; + } - tensorflow::shape_inference::DimensionHandle r; - if (op == "Sub") - TF_RETURN_IF_ERROR(ic->Subtract(da, db, &r)); - else if (op == "Add") - TF_RETURN_IF_ERROR(ic->Add(da, db, &r)); - else if (op == "Mul") - TF_RETURN_IF_ERROR(ic->Multiply(da, db, &r)); - else - TF_RETURN_IF_ERROR( - ic->Divide(da, db, /*evenly_divisible=*/false, &r)); - out_elems.push_back(r); + tensorflow::shape_inference::DimensionHandle r; + absl::Status status; + if (op == "Sub") + status = ic->Subtract(da, db, &r); + else if (op == "Add") + status = ic->Add(da, db, &r); + else if (op == "Mul") + status = ic->Multiply(da, db, &r); + else + status = + ic->Divide(da, db, /*evenly_divisible=*/false, &r); + if (!status.ok()) { + symbolic_propagation_succeeded = false; + break; + } + out_elems.push_back(r); + } + if (symbolic_propagation_succeeded) { + c->output_tensors_as_shapes.resize(1); + c->output_tensors_as_shapes[0] = ic->MakeShape(out_elems); + return absl::OkStatus(); + } } - c->output_tensors_as_shapes.resize(1); - c->output_tensors_as_shapes[0] = ic->MakeShape(out_elems); - // @TODO: Check if we need to do anything with output_tensor_protos. - // S.t c->output_tensor_protos[0] = nullptr; - return absl::OkStatus(); } if (IsConstant(node)) { @@ -1812,6 +1830,19 @@ class SymbolicShapeRefiner { } } else if (IsSize(node)) { DimensionHandle size = ic->NumElements(ic->input(0)); + if (enable_dynamic_value_inference_ && !ic->ValueKnown(size) && + ic->RankKnown(ic->input(0))) { + size = ic->MakeDim(1); + for (int i = 0; i < ic->Rank(ic->input(0)); ++i) { + TF_RETURN_IF_ERROR( + ic->Multiply(size, ic->Dim(ic->input(0), i), &size)); + } + } + if (enable_dynamic_value_inference_ && + (ic->ValueKnown(size) || ic->GetDimExpr(size) != nullptr)) { + c->output_tensors_as_shapes.resize(1); + c->output_tensors_as_shapes[0] = ic->MakeShape({size}); + } if (ic->ValueKnown(size)) { // Propagate size value. int64_t sz = ic->Value(size); @@ -1832,6 +1863,48 @@ class SymbolicShapeRefiner { c->output_tensor_protos[0] = &const_tensors_to_propagate_.back(); } } + } else if (enable_dynamic_value_inference_ && op == "Range") { + auto scalar_int_value = [&](int input, int64_t* value) { + const Tensor* tensor = ic->input_tensor(input); + if (tensor == nullptr || tensor->dims() != 0) { + return false; + } + if (tensor->dtype() == DT_INT32) { + *value = tensor->scalar()(); + return true; + } + if (tensor->dtype() == DT_INT64) { + *value = tensor->scalar()(); + return true; + } + return false; + }; + + int64_t start; + int64_t delta; + const bool has_positive_constant_delta = + scalar_int_value(0, &start) && scalar_int_value(2, &delta) && + start == 0 && delta > 0; + if (has_positive_constant_delta && + c->input_tensors_as_shapes_to_propagate.size() > 1) { + const ShapeHandle& limit = + c->input_tensors_as_shapes_to_propagate[1]; + if (ic->RankKnown(limit) && ic->Rank(limit) >= 1) { + DimensionHandle length = ic->Dim(limit, 0); + if (ic->ValueKnown(length) || ic->GetDimExpr(length) != nullptr) { + DimensionHandle range_length = length; + if (delta > 1) { + DimensionHandle adjusted_length; + TF_RETURN_IF_ERROR( + ic->Add(length, delta - 1, &adjusted_length)); + TF_RETURN_IF_ERROR(ic->Divide( + adjusted_length, delta, /*evenly_divisible=*/false, + &range_length)); + } + ic->set_output(0, ic->Vector(range_length)); + } + } + } } else if (IsShape(node)) { c->output_tensors_as_shapes.resize(1); c->output_tensors_as_shapes[0] = c->inference_context->input(0); @@ -1998,8 +2071,12 @@ class SymbolicShapeRefiner { // possible. const ShapeHandle& shape_handle = c->input_tensors_as_shapes_to_propagate[i]; - if (ic->RankKnown(shape_handle) && ic->Rank(shape_handle) >= 1 && - ic->ValueKnown(ic->Dim(shape_handle, 0))) { + const bool has_value = + ic->RankKnown(shape_handle) && ic->Rank(shape_handle) >= 1 && + (ic->ValueKnown(ic->Dim(shape_handle, 0)) || + (enable_dynamic_value_inference_ && + ic->GetDimExpr(ic->Dim(shape_handle, 0)) != nullptr)); + if (has_value) { dims.push_back(ic->Dim(shape_handle, 0)); } else { // This is not from Const, but as it shouldn'be used as symbolic @@ -2378,6 +2455,7 @@ class SymbolicShapeRefiner { // For more aggressive shape and value inference. bool aggressive_shape_inference_; + bool enable_dynamic_value_inference_; ResourceMgr resource_mgr_; }; @@ -2421,7 +2499,8 @@ class SymbolicShapeManager { shape_inference::DimensionHandle dim = InferenceContext::DimKnownRank(actual_shape, j); int64_t d = dims_.GetMergedValue(dim); - auto* out_dim = properties->mutable_shape()->add_dim(); + TensorShapeProto* output_shape = properties->mutable_shape(); + auto* out_dim = output_shape->add_dim(); out_dim->set_size(d < 0 ? -1 : d); void* root = dims_.RootId(dim); DimExpr* expr = nullptr; @@ -2430,9 +2509,15 @@ class SymbolicShapeManager { } else { expr = ExprForDim(dim); } + ExpressionProto* output_expr = output_shape->add_expressions(); if (expr != nullptr) { - expr->ToProto(out_dim->mutable_expr()); + DimExprToProto(*expr, output_expr); // TODO: Apply simplification? + } else if (d >= 0) { + output_expr->set_constant_value(d); + } else { + DimExprToProto(DimExpr::Unknown(xla::kMissingExpressionSentinel), + output_expr); } } } @@ -2464,7 +2549,7 @@ class SymbolicShapeManager { // Get the variable ID from an expression, or -1 if not a variable. static int32_t GetVarId(const DimExpr* e) { if (!e || e->kind() != DimExpr::Kind::kVariable) return -1; - return static_cast(e)->id(); + return static_cast(e->get())->get_id(); } static bool IsConst(const DimExpr* e) { @@ -2478,7 +2563,7 @@ class SymbolicShapeManager { static bool IsPlaceHolder(const DimExpr* e) { if (!e) return false; if (e->kind() != DimExpr::Kind::kVariable) return false; - return static_cast(e)->id() < 0; + return static_cast(e->get())->get_id() < 0; } static bool IsCompound(const DimExpr* e) { @@ -2488,6 +2573,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; @@ -2534,7 +2622,7 @@ class SymbolicShapeManager { if (it != const_exprs_.end()) { return it->second.get(); } - auto expr = DimExpr::Cons(value); + auto expr = std::make_unique(DimExpr::Const(value)); DimExpr* expr_ptr = expr.get(); const_exprs_.emplace(value, std::move(expr)); return expr_ptr; @@ -2979,7 +3067,8 @@ absl::Status GraphProperties::UpdateEnqueue( absl::Status GraphProperties::InferStatically( bool assume_valid_feeds, bool aggressive_shape_inference, - bool include_input_tensor_values, bool include_output_tensor_values) { + bool include_input_tensor_values, bool include_output_tensor_values, + bool enable_dynamic_value_inference) { FunctionLibraryDefinition function_library(OpRegistry::Global(), item_.graph.library()); absl::flat_hash_map> fed_ports; @@ -3068,7 +3157,8 @@ absl::Status GraphProperties::InferStatically( // Heap-allocate SymbolicShapeRefiner in order to not consume a large amount // of stack space. auto refiner = std::make_unique( - graph_view, fed_ports, aggressive_shape_inference); + graph_view, fed_ports, aggressive_shape_inference, + enable_dynamic_value_inference); TopoQueue new_shapes(topo_order); // Also seed the propagation of shapes in the fanout of primary inputs. diff --git a/tensorflow/core/grappler/costs/graph_properties.h b/tensorflow/core/grappler/costs/graph_properties.h index 1d9575e1e5c805..2d3a51cd2975f0 100644 --- a/tensorflow/core/grappler/costs/graph_properties.h +++ b/tensorflow/core/grappler/costs/graph_properties.h @@ -99,7 +99,8 @@ class GraphProperties { absl::Status InferStatically(bool assume_valid_feeds, bool aggressive_shape_inference, bool include_input_tensor_values, - bool include_output_tensor_values); + bool include_output_tensor_values, + bool enable_dynamic_value_inference = false); absl::Status InferStatically(bool assume_valid_feeds, bool aggressive_shape_inference, bool include_tensor_values) { diff --git a/tensorflow/core/grappler/costs/graph_properties_test.cc b/tensorflow/core/grappler/costs/graph_properties_test.cc index 35dc994092acb4..b24ad629910909 100644 --- a/tensorflow/core/grappler/costs/graph_properties_test.cc +++ b/tensorflow/core/grappler/costs/graph_properties_test.cc @@ -43,12 +43,22 @@ namespace tensorflow { namespace grappler { namespace { -std::string ShapeDimExprDebugString(const TensorShapeProto& shape, int dim) { - if (dim >= shape.expressions_size()) { - return ""; +DimExpr ShapeDimExpr(const TensorShapeProto& shape, int dim) { + if (dim < shape.expressions_size()) { + return DimExprFromProto(shape.expressions(dim)); } - auto expr = DimExpr::FromProto(shape.expressions(dim)); - return expr ? expr->DebugString() : ""; + return DimExpr(); +} + +bool DimExprEqual(const DimExpr& lhs_expr, const DimExpr& rhs_expr) { + if (!lhs_expr || !rhs_expr) return !lhs_expr && !rhs_expr; + return lhs_expr == rhs_expr; +} + +bool ShapeDimExprEqual(const TensorShapeProto& lhs, int lhs_dim, + const TensorShapeProto& rhs, int rhs_dim) { + return DimExprEqual(ShapeDimExpr(lhs, lhs_dim), + ShapeDimExpr(rhs, rhs_dim)); } using shape_inference::InferenceContext; @@ -785,9 +795,11 @@ TEST_F(GraphPropertiesTest, WhileLoop) { // since we concatenated along the batch dim. auto shape_in = properties.GetOutputProperties("ones").at(0).shape(); auto shape_out = properties.GetOutputProperties("while/Exit_1").at(0).shape(); - EXPECT_GE(-2, shape_in.dim(0).size()); - EXPECT_GE(-2, shape_out.dim(0).size()); - EXPECT_NE(shape_in.dim(0).size(), shape_out.dim(0).size()); + EXPECT_EQ(-1, shape_in.dim(0).size()); + EXPECT_EQ(-1, shape_out.dim(0).size()); + EXPECT_TRUE(ShapeDimExpr(shape_in, 0)); + EXPECT_TRUE(ShapeDimExpr(shape_out, 0)); + EXPECT_FALSE(ShapeDimExprEqual(shape_in, 0, shape_out, 0)); } TEST_F(GraphPropertiesTest, NestedLoop) { @@ -1978,10 +1990,10 @@ TEST_F(GraphPropertiesTest, SymbolicShapes) { const auto shape_c = properties.GetOutputProperties("c").at(0).shape(); EXPECT_EQ(2, shape_a.dim_size()); EXPECT_EQ(shape_a.dim_size(), shape_c.dim_size()); - EXPECT_GE(-2, shape_a.dim(0).size()); - EXPECT_EQ(shape_a.dim(0).size(), shape_c.dim(0).size()); - EXPECT_GE(-2, shape_a.dim(1).size()); - EXPECT_EQ(shape_a.dim(1).size(), shape_c.dim(1).size()); + EXPECT_EQ(-1, shape_a.dim(0).size()); + EXPECT_TRUE(ShapeDimExprEqual(shape_a, 0, shape_c, 0)); + EXPECT_EQ(-1, shape_a.dim(1).size()); + EXPECT_TRUE(ShapeDimExprEqual(shape_a, 1, shape_c, 1)); PartialTensorShape shape(shape_a); EXPECT_FALSE(shape.IsFullyDefined()); @@ -1991,29 +2003,29 @@ TEST_F(GraphPropertiesTest, SymbolicShapes) { const auto shape_d = properties.GetOutputProperties("d").at(0).shape(); EXPECT_EQ(1, shape_b.dim_size()); EXPECT_EQ(shape_b.dim_size(), shape_d.dim_size()); - EXPECT_GE(-2, shape_b.dim(0).size()); - EXPECT_NE(shape_a.dim(0).size(), shape_b.dim(0).size()); - EXPECT_EQ(shape_b.dim(0).size(), shape_d.dim(0).size()); + EXPECT_EQ(-1, shape_b.dim(0).size()); + EXPECT_FALSE(ShapeDimExprEqual(shape_a, 0, shape_b, 0)); + EXPECT_TRUE(ShapeDimExprEqual(shape_b, 0, shape_d, 0)); const auto shape_e = properties.GetOutputProperties("e").at(0).shape(); ASSERT_EQ(2, shape_e.dim_size()); - EXPECT_EQ(shape_e.dim(0).size(), shape_c.dim(0).size()); - EXPECT_NE(shape_e.dim(1).size(), shape_c.dim(1).size()); - EXPECT_NE(shape_e.dim(0).size(), shape_d.dim(0).size()); + EXPECT_TRUE(ShapeDimExprEqual(shape_e, 0, shape_c, 0)); + EXPECT_TRUE(ShapeDimExprEqual(shape_e, 1, shape_c, 1)); + EXPECT_FALSE(ShapeDimExprEqual(shape_e, 0, shape_d, 0)); const auto shape_f = properties.GetOutputProperties("f").at(0).shape(); ASSERT_EQ(2, shape_f.dim_size()); - EXPECT_EQ(shape_f.dim(0).size(), shape_a.dim(0).size()); - EXPECT_EQ(shape_f.dim(1).size(), shape_a.dim(1).size()); + EXPECT_TRUE(ShapeDimExprEqual(shape_f, 0, shape_a, 0)); + EXPECT_TRUE(ShapeDimExprEqual(shape_f, 1, shape_a, 1)); const auto shape_h = properties.GetOutputProperties("h").at(0).shape(); ASSERT_EQ(2, shape_f.dim_size()); - EXPECT_EQ(shape_h.dim(0).size(), shape_c.dim(0).size()); - EXPECT_EQ(shape_h.dim(1).size(), shape_c.dim(1).size()); + EXPECT_TRUE(ShapeDimExprEqual(shape_h, 0, shape_c, 0)); + EXPECT_TRUE(ShapeDimExprEqual(shape_h, 1, shape_c, 1)); const auto shape_j = properties.GetOutputProperties("j").at(0).shape(); ASSERT_EQ(1, shape_j.dim_size()); - EXPECT_EQ(shape_j.dim(0).size(), shape_a.dim(1).size()); + EXPECT_TRUE(ShapeDimExprEqual(shape_j, 0, shape_a, 1)); } TEST_F(GraphPropertiesTest, DoNotValidateColocationConstraints) { @@ -2176,6 +2188,74 @@ TEST_F(GraphPropertiesTest, StridedSlicesOfShapes) { EXPECT_EQ(shape_a.dim(1).size(), shape_o2.dim(0).size()); } +TEST_F(GraphPropertiesTest, SizeContentsPropagateToFillOutput) { + tensorflow::Scope scope = tensorflow::Scope::NewRootScope(); + Output input = ops::Placeholder( + scope.WithOpName("input"), DT_FLOAT, + ops::Placeholder::Shape(PartialTensorShape({-1, 24}))); + Output input_size = ops::Size(scope.WithOpName("input_size"), input); + Output fill_shape = + ops::Stack(scope.WithOpName("fill_shape"), {input_size}); + Output zero = ops::Const(scope.WithOpName("zero"), 0.0f, {}); + Output filled = ops::Fill(scope.WithOpName("filled"), fill_shape, zero); + + GrapplerItem item; + TF_ASSERT_OK(scope.ToGraphDef(&item.graph)); + + GraphProperties properties(item); + TF_ASSERT_OK(properties.InferStatically( + /*assume_valid_feeds=*/false, + /*aggressive_shape_inference=*/false, + /*include_input_tensor_values=*/false, + /*include_output_tensor_values=*/false, + /*enable_dynamic_value_inference=*/true)); + + const TensorShapeProto& inferred_input_shape = + properties.GetOutputProperties("input").at(0).shape(); + const TensorShapeProto& inferred_fill_shape = + properties.GetOutputProperties("filled").at(0).shape(); + + ASSERT_EQ(1, inferred_fill_shape.dim_size()); + ASSERT_EQ(2, inferred_input_shape.expressions_size()); + DimExpr expected = + DimExprFromProto(inferred_input_shape.expressions(0)) * 24; + EXPECT_TRUE(DimExprEqual(expected, ShapeDimExpr(inferred_fill_shape, 0))); +} + +TEST_F(GraphPropertiesTest, SizeContentsPropagateToRangeOutput) { + tensorflow::Scope scope = tensorflow::Scope::NewRootScope(); + Output input = ops::Placeholder( + scope.WithOpName("input"), DT_FLOAT, + ops::Placeholder::Shape(PartialTensorShape({-1, 1}))); + Output input_size = ops::Size(scope.WithOpName("input_size"), input); + Output zero = ops::Const(scope.WithOpName("zero"), 0, {}); + Output three = ops::Const(scope.WithOpName("three"), 3, {}); + Output range = + ops::Range(scope.WithOpName("range"), zero, input_size, three); + + GrapplerItem item; + TF_ASSERT_OK(scope.ToGraphDef(&item.graph)); + + GraphProperties properties(item); + TF_ASSERT_OK(properties.InferStatically( + /*assume_valid_feeds=*/false, + /*aggressive_shape_inference=*/false, + /*include_input_tensor_values=*/false, + /*include_output_tensor_values=*/false, + /*enable_dynamic_value_inference=*/true)); + + const TensorShapeProto& inferred_input_shape = + properties.GetOutputProperties("input").at(0).shape(); + const TensorShapeProto& inferred_range_shape = + properties.GetOutputProperties("range").at(0).shape(); + + ASSERT_EQ(1, inferred_range_shape.dim_size()); + ASSERT_EQ(2, inferred_input_shape.expressions_size()); + DimExpr expected = + (DimExprFromProto(inferred_input_shape.expressions(0)) + 2) / 3; + EXPECT_TRUE(DimExprEqual(expected, ShapeDimExpr(inferred_range_shape, 0))); +} + TEST_F(GraphPropertiesTest, StridedSliceOfShapeWithShrinkAxisMask) { tensorflow::Scope scope = tensorflow::Scope::NewRootScope(); Output placeholder = @@ -2343,20 +2423,17 @@ TEST_F(GraphPropertiesTest, ShapeTensorContentsThroughGatherProdAndUnpack) { ASSERT_EQ(3, input_shape.dim_size()); ASSERT_EQ(2, gather_reshape_shape.dim_size()); ASSERT_EQ(2, unpack_reshape_shape.dim_size()); - ASSERT_GT(input_shape.expressions_size(), 0); - ASSERT_GT(gather_reshape_shape.expressions_size(), 0); - ASSERT_GT(unpack_reshape_shape.expressions_size(), 0); - - auto expected_gather_dim0 = - std::make_unique(DimExpr::FromProto(input_shape.expressions(0)) - .release(), - new Constant(26)); - EXPECT_EQ(expected_gather_dim0->DebugString(), - ShapeDimExprDebugString(gather_reshape_shape, 0)); + ASSERT_TRUE(ShapeDimExpr(input_shape, 0)); + ASSERT_TRUE(ShapeDimExpr(gather_reshape_shape, 0)); + ASSERT_TRUE(ShapeDimExpr(unpack_reshape_shape, 0)); + + DimExpr input_expr = ShapeDimExpr(input_shape, 0); + auto expected_gather_dim0 = input_expr * DimExpr::Const(26); + EXPECT_TRUE(DimExprEqual(expected_gather_dim0, + ShapeDimExpr(gather_reshape_shape, 0))); EXPECT_EQ(8, gather_reshape_shape.dim(1).size()); - EXPECT_EQ(ShapeDimExprDebugString(input_shape, 0), - ShapeDimExprDebugString(unpack_reshape_shape, 0)); + EXPECT_TRUE(ShapeDimExprEqual(input_shape, 0, unpack_reshape_shape, 0)); EXPECT_EQ(208, unpack_reshape_shape.dim(1).size()); } @@ -2404,6 +2481,37 @@ TEST_F(GraphPropertiesTest, ValuePropagationThroughArithmeticOps) { ExpectTensorValues({20, 24}, c_plus_b_plus_2a_prop.value()); } +TEST_F(GraphPropertiesTest, + DynamicArithmeticFallsBackToTensorValueInference) { + tensorflow::Scope scope = tensorflow::Scope::NewRootScope(); + Output one = ops::Const(scope.WithOpName("one"), {1}, {1}); + Output two = ops::Const(scope.WithOpName("two"), {2}, {1}); + Output unknown = ops::Const(scope.WithOpName("unknown"), {-1}, {1}); + Output negative = ops::Sub(scope.WithOpName("negative"), one, two); + Output unknown_plus_one = + ops::Add(scope.WithOpName("unknown_plus_one"), unknown, one); + + GrapplerItem item; + TF_ASSERT_OK(scope.ToGraphDef(&item.graph)); + GraphProperties properties(item); + TF_ASSERT_OK(properties.InferStatically( + /*assume_valid_feeds=*/false, + /*aggressive_shape_inference=*/true, + /*include_input_tensor_values=*/true, + /*include_output_tensor_values=*/true, + /*enable_dynamic_value_inference=*/true)); + + const auto& negative_prop = + properties.GetOutputProperties("negative").at(0); + ASSERT_TRUE(negative_prop.has_value()); + ExpectTensorValues({-1}, negative_prop.value()); + + const auto& unknown_plus_one_prop = + properties.GetOutputProperties("unknown_plus_one").at(0); + ASSERT_TRUE(unknown_plus_one_prop.has_value()); + ExpectTensorValues({0}, unknown_plus_one_prop.value()); +} + TEST_F(GraphPropertiesTest, ShapeAnnotation) { GrapplerItem item; TF_ASSERT_OK(NodeDefBuilder("Input", "Placeholder") diff --git a/third_party/xla/xla/hlo/builder/lib/constants.cc b/third_party/xla/xla/hlo/builder/lib/constants.cc index acfa2fe0b66e2c..da6a97f4088520 100644 --- a/third_party/xla/xla/hlo/builder/lib/constants.cc +++ b/third_party/xla/xla/hlo/builder/lib/constants.cc @@ -33,7 +33,8 @@ XlaOp Zero(XlaBuilder* builder, PrimitiveType type) { } XlaOp Zeros(XlaBuilder* builder, const Shape& shape) { - return Broadcast(Zero(builder, shape.element_type()), shape.dimensions()); + return Broadcast(Zero(builder, shape.element_type()), shape.dimensions(), + shape.expressions()); } XlaOp ZerosLike(XlaOp prototype) { 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/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/cpu/BUILD b/third_party/xla/xla/service/cpu/BUILD index e5132538349bbe..850eb7e7239f4c 100644 --- a/third_party/xla/xla/service/cpu/BUILD +++ b/third_party/xla/xla/service/cpu/BUILD @@ -690,6 +690,7 @@ cc_library( "//xla/service:custom_call_status_internal", "//xla/service:executable", "//xla/service:hlo_execution_profile", + "//xla/service:hlo_profile_printer", "//xla/service:hlo_profile_printer_data_cc", "//xla/service:hlo_value", "//xla/service:maybe_owning_device_memory", @@ -1123,6 +1124,7 @@ cc_library( "//xla/stream_executor:stream_executor_h", "//xla/tsl/concurrency:async_value", "//xla/tsl/platform:errors", + "//xla/tsl/platform:env_time", "//xla/tsl/platform:logging", "//xla/tsl/platform:status", "@com_google_absl//absl/algorithm:container", diff --git a/third_party/xla/xla/service/cpu/cpu_compiler.cc b/third_party/xla/xla/service/cpu/cpu_compiler.cc index 042e7e5b33a0f7..b31b1b2e3cad24 100644 --- a/third_party/xla/xla/service/cpu/cpu_compiler.cc +++ b/third_party/xla/xla/service/cpu/cpu_compiler.cc @@ -1063,6 +1063,25 @@ absl::Status CreateHloProfilingArtifacts( return absl::OkStatus(); } +bool ShouldCreateHloProfilingArtifacts(const HloModule& module) { + if (!module.config().hlo_profiling_enabled()) { + return false; + } + + // The thunk runtime launches host kernels through XLA_CPU_KernelCallFrame, + // which currently has no profile-counters field. Emitting profiling loads and + // stores in this mode can make generated code dereference a null counter + // pointer at run time. + if (module.config().debug_options().xla_cpu_use_thunk_runtime()) { + LOG(WARNING) << "--xla_hlo_profile is not supported by XLA:CPU thunk " + "runtime; disabling HLO profiling for module " + << module.name(); + return false; + } + + return true; +} + } // namespace absl::StatusOr> CpuCompiler::RunHloPasses( @@ -1451,7 +1470,7 @@ CpuCompiler::CompileCpuExecutable(std::unique_ptr module) { computation_to_profile_idx; std::unique_ptr hlo_profile_index_map; std::unique_ptr hlo_profile_printer_data; - if (module->config().hlo_profiling_enabled()) { + if (ShouldCreateHloProfilingArtifacts(*module)) { TF_RETURN_IF_ERROR(CreateHloProfilingArtifacts( *module, &instruction_to_profile_idx, &computation_to_profile_idx, &hlo_profile_index_map, &hlo_profile_printer_data)); @@ -2003,7 +2022,7 @@ CpuCompiler::CompileAheadOfTimeLegacy( std::unique_ptr hlo_profile_index_map; std::unique_ptr hlo_profile_printer_data; - if (module->config().hlo_profiling_enabled()) { + if (ShouldCreateHloProfilingArtifacts(*module)) { TF_RETURN_IF_ERROR(CreateHloProfilingArtifacts( *module, &instruction_to_profile_idx, &computation_to_profile_idx, &hlo_profile_index_map, &hlo_profile_printer_data)); @@ -2161,7 +2180,7 @@ CpuCompiler::CompileAheadOfTimeThunks( computation_to_profile_idx; std::unique_ptr hlo_profile_index_map; std::unique_ptr hlo_profile_printer_data; - if (module->config().hlo_profiling_enabled()) { + if (ShouldCreateHloProfilingArtifacts(*module)) { TF_RETURN_IF_ERROR(CreateHloProfilingArtifacts( *module, &instruction_to_profile_idx, &computation_to_profile_idx, &hlo_profile_index_map, &hlo_profile_printer_data)); @@ -2410,7 +2429,7 @@ CpuCompiler::CompileAheadOfTimeThunks( cpu_executable->thunks().thunk_sequence(); std::unique_ptr executable_hlo_profile_printer_data = - cpu_executable->module().config().hlo_profiling_enabled() + cpu_executable->hlo_profiling_enabled() ? std::make_unique( cpu_executable->hlo_profile_printer_data()) : nullptr; diff --git a/third_party/xla/xla/service/cpu/cpu_executable.cc b/third_party/xla/xla/service/cpu/cpu_executable.cc index de9ffc2b78eb80..6ae5465b5b79bb 100644 --- a/third_party/xla/xla/service/cpu/cpu_executable.cc +++ b/third_party/xla/xla/service/cpu/cpu_executable.cc @@ -54,6 +54,7 @@ limitations under the License. #include "xla/service/custom_call_status_internal.h" #include "xla/service/executable.h" #include "xla/service/hlo_execution_profile.h" +#include "xla/service/hlo_profile_printer.h" #include "xla/service/hlo_profile_printer_data.pb.h" #include "xla/service/hlo_value.h" #include "xla/service/maybe_owning_device_memory.h" @@ -76,6 +77,20 @@ limitations under the License. namespace xla { namespace cpu { +namespace { + +float GetClockRateGhz(const ExecutableRunOptions* run_options) { + if (run_options->stream() == nullptr || + run_options->stream()->parent() == nullptr) { + return 1.0f; + } + float clock_rate_ghz = + run_options->stream()->parent()->GetDeviceDescription().clock_rate_ghz(); + return clock_rate_ghz > 0 ? clock_rate_ghz : 1.0f; +} + +} // namespace + absl::StatusOr> CpuExecutable::Create( std::unique_ptr function_library, std::unique_ptr assignment, @@ -246,8 +261,16 @@ absl::Status CpuExecutable::ExecuteComputeFunction( absl::Span buffers) { uint64_t start_micros = tsl::Env::Default()->NowMicros(); - size_t profile_counters_size = 0; + std::vector profile_counter_storage; + if (hlo_profiling_enabled()) { + profile_counter_storage.assign( + hlo_profile_printer_data().profile_counters_size(), 0); + } + size_t profile_counters_size = profile_counter_storage.size(); int64_t* profile_counters = nullptr; + if (!profile_counter_storage.empty()) { + profile_counters = profile_counter_storage.data(); + } // Call the computation function following the calling convention. See the // definition of 'ComputeFunctionType' for the details of the calling @@ -286,6 +309,15 @@ absl::Status CpuExecutable::ExecuteComputeFunction( compute_function_(nullptr, run_options, nullptr, buffer_pointers.data(), &status, profile_counters); record_profile(); + if (profile_counters != nullptr) { + std::string hlo_profile = + PrintHloProfile(hlo_profile_printer_data(), profile_counters, + GetClockRateGhz(run_options)); + if (!hlo_profile.empty()) { + LOG(INFO) << "XLA:CPU HLO profile for " << module().name() << "\n" + << hlo_profile; + } + } std::optional error_message = CustomCallStatusGetMessage(&status); if (error_message) { diff --git a/third_party/xla/xla/service/cpu/cpu_runtime.cc b/third_party/xla/xla/service/cpu/cpu_runtime.cc index 7caf9c43b1119b..20c482f663b6c6 100644 --- a/third_party/xla/xla/service/cpu/cpu_runtime.cc +++ b/third_party/xla/xla/service/cpu/cpu_runtime.cc @@ -60,6 +60,7 @@ limitations under the License. #include "xla/stream_executor/stream_executor.h" #include "xla/tsl/concurrency/async_value_ref.h" #include "xla/tsl/platform/errors.h" +#include "xla/tsl/platform/env_time.h" #include "xla/tsl/platform/logging.h" #include "xla/tsl/platform/status.h" #include "xla/util.h" @@ -166,6 +167,8 @@ extern const char* const kParallelForkJoinSymbolName = "__xla_cpu_runtime_ParallelForkJoin"; extern const char* const kPrintfToStderrSymbolName = "__xla_cpu_runtime_PrintfToStderr"; +extern const char* const kReadCycleCounterSymbolName = + "__xla_cpu_runtime_ReadCycleCounter"; extern const char* const kStatusIsSuccessSymbolName = "__xla_cpu_runtime_StatusIsSuccess"; extern const char* const kKeyValueSortSymbolName = @@ -620,6 +623,11 @@ ABSL_ATTRIBUTE_NO_SANITIZE_MEMORY int __xla_cpu_runtime_PrintfToStderr( return result; } +ABSL_ATTRIBUTE_NO_SANITIZE_MEMORY uint64_t +__xla_cpu_runtime_ReadCycleCounter() { + return tsl::EnvTime::NowNanos(); +} + ABSL_ATTRIBUTE_NO_SANITIZE_MEMORY int64_t __xla_cpu_runtime_TracingStart( const void* /* ExecutableRunOptions* run_options_ptr*/, const char* name, const char* hlo_module, int64_t program_id) { diff --git a/third_party/xla/xla/service/cpu/cpu_runtime.h b/third_party/xla/xla/service/cpu/cpu_runtime.h index 71e27ea600ee28..de1ea93a855a35 100644 --- a/third_party/xla/xla/service/cpu/cpu_runtime.h +++ b/third_party/xla/xla/service/cpu/cpu_runtime.h @@ -79,6 +79,7 @@ extern const char* const kAcquireOutfeedBufferForPopulationSymbolName; extern const char* const kReleaseOutfeedBufferAfterPopulationSymbolName; extern const char* const kParallelForkJoinSymbolName; extern const char* const kPrintfToStderrSymbolName; +extern const char* const kReadCycleCounterSymbolName; extern const char* const kStatusIsSuccessSymbolName; extern const char* const kKeyValueSortSymbolName; extern const char* const kTopKF32SymbolName; @@ -115,6 +116,7 @@ int GetDeviceOrdinal(const xla::ExecutableRunOptions* run_options); extern "C" { extern int __xla_cpu_runtime_PrintfToStderr(const char* format, ...); +extern uint64_t __xla_cpu_runtime_ReadCycleCounter(); extern int64_t __xla_cpu_runtime_TracingStart( const void* /* xla::ExecutableRunOptions* */ run_options_ptr, diff --git a/third_party/xla/xla/service/cpu/ir_emitter.cc b/third_party/xla/xla/service/cpu/ir_emitter.cc index ae40e76e8e2ce2..bdf1bc2e60f4df 100644 --- a/third_party/xla/xla/service/cpu/ir_emitter.cc +++ b/third_party/xla/xla/service/cpu/ir_emitter.cc @@ -3715,10 +3715,26 @@ void IrEmitter::ProfilingState::UpdateProfileCounter(llvm::IRBuilderBase* b, llvm::Value* prof_counter, llvm::Value* cycle_end, llvm::Value* cycle_start) { + llvm::Value* profile_counters = prof_counter; + if (auto* gep = llvm::dyn_cast(prof_counter)) { + profile_counters = gep->getPointerOperand(); + } + + auto* profile_counters_type = + llvm::cast(profile_counters->getType()); + llvm::Value* has_profile_counters = b->CreateICmpNE( + profile_counters, llvm::ConstantPointerNull::get(profile_counters_type), + "has_profile_counters"); + + llvm::AllocaInst* dummy_counter = llvm_ir::EmitAllocaAtFunctionEntry( + b->getInt64Ty(), "dummy_profile_counter", b); + b->CreateStore(b->getInt64(0), dummy_counter); + prof_counter = + b->CreateSelect(has_profile_counters, prof_counter, dummy_counter); + auto* cycle_diff = b->CreateSub(cycle_end, cycle_start); llvm::LoadInst* old_cycle_count = b->CreateLoad( - llvm::cast(prof_counter)->getSourceElementType(), - prof_counter, "old_cycle_count"); + b->getInt64Ty(), prof_counter, "old_cycle_count"); auto* new_cycle_count = b->CreateAdd(cycle_diff, old_cycle_count, "new_cycle_count"); b->CreateStore(new_cycle_count, prof_counter); @@ -3728,10 +3744,17 @@ llvm::Value* IrEmitter::ProfilingState::ReadCycleCounter( llvm::IRBuilderBase* b) { llvm::Module* module = b->GetInsertBlock()->getModule(); if (!use_rdtscp_) { - llvm::Function* func_llvm_readcyclecounter = - llvm::Intrinsic::getOrInsertDeclaration( - module, llvm::Intrinsic::readcyclecounter); - return b->CreateCall(func_llvm_readcyclecounter); + llvm::FunctionType* fn_type = + llvm::FunctionType::get(b->getInt64Ty(), /*isVarArg=*/false); + llvm::FunctionCallee read_cycle_counter_func = + module->getOrInsertFunction(runtime::kReadCycleCounterSymbolName, + fn_type); + if (auto* fn = + llvm::dyn_cast(read_cycle_counter_func.getCallee())) { + fn->setCallingConv(llvm::CallingConv::C); + fn->setDoesNotThrow(); + } + return b->CreateCall(read_cycle_counter_func); } llvm::Function* func_llvm_x86_rdtscp = llvm::Intrinsic::getOrInsertDeclaration(module, diff --git a/third_party/xla/xla/service/cpu/runtime_symbol_generator.cc b/third_party/xla/xla/service/cpu/runtime_symbol_generator.cc index 87aca6c386751a..c7e95d8ed02aa1 100644 --- a/third_party/xla/xla/service/cpu/runtime_symbol_generator.cc +++ b/third_party/xla/xla/service/cpu/runtime_symbol_generator.cc @@ -201,6 +201,7 @@ static bool RegisterKnownJITSymbols() { REGISTER_CPU_RUNTIME_SYMBOL(EigenSingleThreadedMatMulU8); REGISTER_CPU_RUNTIME_SYMBOL(ParallelForkJoin); REGISTER_CPU_RUNTIME_SYMBOL(PrintfToStderr); + REGISTER_CPU_RUNTIME_SYMBOL(ReadCycleCounter); REGISTER_CPU_RUNTIME_SYMBOL(ReleaseInfeedBufferAfterDequeue); REGISTER_CPU_RUNTIME_SYMBOL(ReleaseOutfeedBufferAfterPopulation); REGISTER_CPU_RUNTIME_SYMBOL(StatusIsSuccess); 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/service/llvm_ir/BUILD b/third_party/xla/xla/service/llvm_ir/BUILD index 2a665c2d77f3b2..7609f7d94ac1e6 100644 --- a/third_party/xla/xla/service/llvm_ir/BUILD +++ b/third_party/xla/xla/service/llvm_ir/BUILD @@ -356,3 +356,17 @@ xla_cc_test( "@llvm-project//llvm:ir_headers", ], ) + +xla_cc_test( + name = "llvm_util_test", + srcs = ["llvm_util_test.cc"], + deps = [ + ":llvm_util", + "//xla:shape_util", + "//xla/tests:xla_internal_test_main", + "@com_google_googletest//:gtest", + "@llvm-project//llvm:Core", + "@llvm-project//llvm:Support", + "@llvm-project//llvm:ir_headers", + ], +) 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..4d6804a506a6b5 100644 --- a/third_party/xla/xla/service/llvm_ir/llvm_util.cc +++ b/third_party/xla/xla/service/llvm_ir/llvm_util.cc @@ -905,7 +905,7 @@ static llvm::Value* EmitExpressionImpl(llvm::IRBuilderBase* b, auto* div_node = static_cast(&expr); llvm::Value* v_lhs = EmitExpressionImpl(b, *div_node->get_lhs()); llvm::Value* v_rhs = EmitExpressionImpl(b, *div_node->get_rhs()); - return b->CreateUDiv(v_lhs, v_rhs, "div_dims"); + return b->CreateSDiv(v_lhs, v_rhs, "div_dims"); } if (expr.kind() == DExpr::Kind::kAdd) { auto* add_node = static_cast(&expr); @@ -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/llvm_ir/llvm_util_test.cc b/third_party/xla/xla/service/llvm_ir/llvm_util_test.cc new file mode 100644 index 00000000000000..fb3a839085a41b --- /dev/null +++ b/third_party/xla/xla/service/llvm_ir/llvm_util_test.cc @@ -0,0 +1,59 @@ +/* Copyright 2026 The OpenXLA Authors. + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +==============================================================================*/ + +#include "xla/service/llvm_ir/llvm_util.h" + +#include +#include "llvm/IR/BasicBlock.h" +#include "llvm/IR/Function.h" +#include "llvm/IR/IRBuilder.h" +#include "llvm/IR/Instructions.h" +#include "llvm/IR/LLVMContext.h" +#include "llvm/IR/Module.h" +#include "llvm/IR/Type.h" +#include "llvm/Support/Casting.h" +#include "xla/shape_expr.h" + +namespace xla { +namespace llvm_ir { +namespace { + +TEST(LlvmUtilTest, DynamicExpressionDivisionIsSigned) { + llvm::LLVMContext context; + llvm::Module module("llvm_util_test", context); + llvm::IRBuilder<> builder(context); + llvm::FunctionType* function_type = llvm::FunctionType::get( + llvm::Type::getVoidTy(context), /*isVarArg=*/false); + llvm::Function* function = llvm::Function::Create( + function_type, llvm::Function::ExternalLinkage, "test", module); + llvm::BasicBlock* block = + llvm::BasicBlock::Create(context, "entry", function); + builder.SetInsertPoint(block); + + llvm::Type* i64_type = builder.getInt64Ty(); + llvm::Value* batch_dim_address = builder.CreateAlloca(i64_type); + builder.CreateLoad(i64_type, batch_dim_address, "bdim_value"); + + DExpr expression = (DExpr::Var(1) - 5) / 2; + llvm::Value* value = EmitExpression(&builder, expression); + + auto* division = llvm::dyn_cast(value); + ASSERT_NE(division, nullptr); + EXPECT_EQ(division->getOpcode(), llvm::Instruction::SDiv); +} + +} // namespace +} // namespace llvm_ir +} // namespace xla diff --git a/third_party/xla/xla/service/shape_inference.cc b/third_party/xla/xla/service/shape_inference.cc index 4226e901326321..ce006f2a132dd3 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,21 +226,40 @@ 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); + const int64_t dilated_window = + window_util::DilatedBound(dim.size(), dim.window_dilation()); + 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 = - window_util::DilatedBound(dim.size(), dim.window_dilation()); - output_dimensions[i] = window_util::StridedBound( 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 = 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 strided_bound_expr = + (padded_dilated_base_expr - dilated_window + 1 + dim.stride() - 1) / + dim.stride(); + output_expressions[i] = + DExpr::Max(strided_bound_expr, DExpr::Const(0)).simplify(); } return ShapeUtil::MakeValidatedShape(element_type, output_dimensions, @@ -2460,11 +2488,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, @@ -3324,7 +3354,7 @@ ShapeInference::InferCollectivePermuteDoneShape(const Shape& operand_shape) { auto new_expr = limit_expr - start_expr + DExpr::Const(stride) - DExpr::Const(1); - expressions.push_back(new_expr / DExpr::Const(stride)); + expressions.push_back((new_expr / DExpr::Const(stride)).simplify()); } std::vector is_dynamic(arg.dimensions_size()); @@ -3900,6 +3930,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 +4167,22 @@ ShapeInference::InferCollectivePermuteDoneShape(const Shape& operand_shape) { on_true.is_dynamic_dimension(dimension) || on_false.is_dynamic_dimension(dimension)); } + const DExpr& on_true_expr = on_true.expressions(dimension); + const DExpr& on_false_expr = on_false.expressions(dimension); + const bool has_dynamic_expression = + (on_true_expr && on_true_expr->is_dynamic()) || + (on_false_expr && on_false_expr->is_dynamic()); + if (has_dynamic_expression) { + if (!on_true_expr || !on_false_expr || + !DynExpr::equal(on_true_expr, on_false_expr)) { + 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_expr); + } } 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..0b6e4bf00650b2 100644 --- a/third_party/xla/xla/service/shape_inference_test.cc +++ b/third_party/xla/xla/service/shape_inference_test.cc @@ -205,6 +205,47 @@ 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, SelectRejectsMissingDynamicOperandExpression) { + 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}); + 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 +489,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 +819,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 +1435,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}); @@ -1369,6 +1531,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}); @@ -1532,6 +1746,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 +1818,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 +1928,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 +2087,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 +2599,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 +2669,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 +3200,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 +3419,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 +3540,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 +3607,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 +3695,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 +3981,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 +4201,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 +5393,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 +5486,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 +5542,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 +5952,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 +5990,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 +6022,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 +6293,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 +6567,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..d74a42175b3380 100644 --- a/third_party/xla/xla/shape_expr.cc +++ b/third_party/xla/xla/shape_expr.cc @@ -23,8 +23,10 @@ limitations under the License. #include #include #include +#include #include #include +#include #include "absl/log/check.h" #include "xla/printer.h" @@ -39,6 +41,131 @@ Constant* AsConstant(DynExpr* expr) { : nullptr; } +std::vector ExpressionChildren(DynExpr* expr) { + CHECK(expr != nullptr); + switch (expr->kind()) { + case DExpr::Kind::kUnknown: + case DExpr::Kind::kConstant: + case DExpr::Kind::kVariable: + return {}; + case DExpr::Kind::kAdd: { + auto* add = static_cast(expr); + return {add->get_lhs(), add->get_rhs()}; + } + case DExpr::Kind::kSub: { + auto* sub = static_cast(expr); + return {sub->get_lhs(), sub->get_rhs()}; + } + case DExpr::Kind::kMul: { + auto* mul = static_cast(expr); + return {mul->get_lhs(), mul->get_rhs()}; + } + case DExpr::Kind::kDiv: { + auto* div = static_cast(expr); + return {div->get_lhs(), div->get_rhs()}; + } + case DExpr::Kind::kMax: { + auto* max = static_cast(expr); + return {max->get_lhs(), max->get_rhs()}; + } + case DExpr::Kind::kGt: { + auto* gt = static_cast(expr); + return {gt->get_lhs(), gt->get_rhs()}; + } + case DExpr::Kind::kSelect: { + auto* select = static_cast(expr); + return {select->get_pred(), select->get_on_true(), + select->get_on_false()}; + } + } + return {}; +} + +DynExpr* FindSmallestCoveringSubexpression(DynExpr* expr) { + CHECK(expr != nullptr); + if (expr->kind() == DExpr::Kind::kVariable) { + return expr; + } + + DynExpr* common_core = nullptr; + for (DynExpr* child : ExpressionChildren(expr)) { + DynExpr* child_core = FindSmallestCoveringSubexpression(child); + if (child_core == nullptr) { + continue; + } + if (common_core == nullptr) { + common_core = child_core; + } else if (!DynExpr::equal(common_core, child_core)) { + return expr; + } + } + return common_core; +} + +std::unique_ptr ReplaceSubexpression(DynExpr* expr, DynExpr* target, + DynExpr* replacement) { + CHECK(expr != nullptr); + CHECK(target != nullptr); + CHECK(replacement != nullptr); + if (DynExpr::equal(expr, target)) { + return replacement->clone(); + } + + switch (expr->kind()) { + case DExpr::Kind::kUnknown: + case DExpr::Kind::kConstant: + case DExpr::Kind::kVariable: + return expr->clone(); + case DExpr::Kind::kAdd: { + auto* add = static_cast(expr); + return std::make_unique( + ReplaceSubexpression(add->get_lhs(), target, replacement).release(), + ReplaceSubexpression(add->get_rhs(), target, replacement).release()); + } + case DExpr::Kind::kSub: { + auto* sub = static_cast(expr); + return std::make_unique( + ReplaceSubexpression(sub->get_lhs(), target, replacement).release(), + ReplaceSubexpression(sub->get_rhs(), target, replacement).release()); + } + case DExpr::Kind::kMul: { + auto* mul = static_cast(expr); + return std::make_unique( + ReplaceSubexpression(mul->get_lhs(), target, replacement).release(), + ReplaceSubexpression(mul->get_rhs(), target, replacement).release()); + } + case DExpr::Kind::kDiv: { + auto* div = static_cast(expr); + return std::make_unique
( + ReplaceSubexpression(div->get_lhs(), target, replacement).release(), + ReplaceSubexpression(div->get_rhs(), target, replacement).release()); + } + case DExpr::Kind::kMax: { + auto* max = static_cast(expr); + return std::make_unique( + ReplaceSubexpression(max->get_lhs(), target, replacement).release(), + ReplaceSubexpression(max->get_rhs(), target, replacement).release()); + } + case DExpr::Kind::kGt: { + auto* gt = static_cast(expr); + return std::make_unique( + ReplaceSubexpression(gt->get_lhs(), target, replacement).release(), + ReplaceSubexpression(gt->get_rhs(), target, replacement).release()); + } + case DExpr::Kind::kSelect: { + auto* select = static_cast(expr); + return std::make_unique( + ReplaceSubexpression(select->get_pred(), target, replacement) + .release(), + ReplaceSubexpression(select->get_on_true(), target, replacement) + .release(), + ReplaceSubexpression(select->get_on_false(), target, replacement) + .release()); + } + } + return expr->clone(); +} + void NormalizeFraction(int64_t* numerator, int64_t* denominator) { CHECK(denominator != nullptr); CHECK(*denominator != 0); @@ -208,10 +335,31 @@ 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; } +bool IsNonNegativeForPositiveVariables(const DynExpr* expr) { + auto affine = ToCanonicalAffine(expr); + if (!affine.has_value()) return false; + + // Dynamic dimension variables are strictly positive. An affine expression + // has a finite lower bound only when every variable coefficient is + // non-negative; evaluate that bound with each variable set to one. + __int128 lower_bound = affine->constant; + for (const auto& [_, coefficient] : affine->coefficients) { + if (coefficient < 0) return false; + lower_bound += coefficient; + } + return lower_bound >= 0; +} + std::unique_ptr BuildScaledVariableTerm(int id, int64_t coefficient) { CHECK(coefficient != 0); if (coefficient == 1) { @@ -231,11 +379,14 @@ std::unique_ptr BuildAffineNumerator(const CanonicalAffineExpr& expr) { } } if (expr.constant != 0 || result == nullptr) { - auto constant_term = std::make_unique(expr.constant); if (result == nullptr) { - result = std::move(constant_term); + result = std::make_unique(expr.constant); + } else if (expr.constant > 0) { + result = std::make_unique(result.release(), + DynExpr::_(expr.constant)); } else { - result = std::make_unique(result.release(), constant_term.release()); + result = std::make_unique(result.release(), + DynExpr::_(-expr.constant)); } } return result; @@ -317,10 +468,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(); @@ -331,8 +496,67 @@ 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 && + IsNonNegativeForPositiveVariables(rhs.get())) { + return rhs; + } + if (r && r->get_val() == 0 && + IsNonNegativeForPositiveVariables(lhs.get())) { + 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->kind() == DExpr::Kind::kUnknown || + rhs->kind() == DExpr::Kind::kUnknown) { + return std::make_unique(); + } + 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 (pred->kind() == DExpr::Kind::kUnknown) { + return std::make_unique(); + } + if (pred->is_constant()) { + return pred->get_val() != 0 ? std::move(on_true) : std::move(on_false); + } + if (on_true->kind() == DExpr::Kind::kUnknown || + on_false->kind() == DExpr::Kind::kUnknown) { + return std::make_unique(); + } + if (*on_true == *on_false) return on_true; + 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) { @@ -340,6 +564,11 @@ std::unique_ptr SimplifyCanonical(const DynExpr* expr) { return std::make_unique(); } if (auto canonical = ToCanonicalAffine(expr); canonical.has_value()) { + if (canonical->IsPureConstant()) { + CHECK_NE(canonical->denominator, 0); + return std::make_unique(canonical->constant / + canonical->denominator); + } return BuildCanonicalExpr(*canonical); } return SimplifyFallback(expr); @@ -347,6 +576,24 @@ std::unique_ptr SimplifyCanonical(const DynExpr* expr) { } // namespace +DExpr DExpr::find_smallest_subexpression_covering_all_variables() const { + CHECK(expr_ != nullptr); + const std::set ids = expr_->get_all_ids(); + CHECK(!ids.empty()); + DynExpr* result = FindSmallestCoveringSubexpression(expr_.get()); + CHECK(result != nullptr); + return DExpr::Adopt(result->clone().release()); +} + +DExpr DExpr::replace_subexpression(const DExpr& target, + const DExpr& replacement) const { + if (expr_ == nullptr) { + return DExpr(); + } + return DExpr(ReplaceSubexpression(expr_.get(), target.get(), + replacement.get())); +} + const DExpr& Shape::MissingExpression() { static const DExpr missing = DExpr::Unknown(kMissingExpressionSentinel); return missing; @@ -358,6 +605,9 @@ DynExpr* operator*(DynExpr& lhs, DynExpr& rhs) { DynExpr* operator*(int64_t k, DynExpr& rhs) { return new Mul(DynExpr::_(k), rhs.clone().release()); } +DynExpr* operator*(DynExpr& lhs, int64_t k) { + return new Mul(lhs.clone().release(), DynExpr::_(k)); +} DynExpr* operator/(DynExpr& lhs, DynExpr& rhs) { return new Div(lhs.clone().release(), rhs.clone().release()); } @@ -387,10 +637,32 @@ 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()); 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() == @@ -443,6 +715,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; } @@ -458,6 +753,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..babd157a56b0b9 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; } @@ -141,6 +151,12 @@ class DExpr { DExpr substitute(int id, const DExpr& value) const { return expr_ == nullptr ? DExpr() : Adopt(expr_->substitute(id, value.get())); } + // Returns the smallest subtree containing every variable in this expression. + // The expression must contain at least one variable. + DExpr find_smallest_subexpression_covering_all_variables() const; + // Replaces every subtree equivalent to `target` with `replacement`. + DExpr replace_subexpression(const DExpr& target, + const DExpr& replacement) const; template friend H AbslHashValue(H h, const DExpr& expr) { @@ -185,8 +201,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(); } @@ -226,7 +241,8 @@ class Constant : public DynExpr { DynExpr* s() override; }; -// var id (int) +// Root variables represent strictly positive dynamic dimensions. Potentially +// signed values must be represented by expressions derived from these roots. class Variable : public DynExpr { int id; @@ -486,13 +502,12 @@ class Div : public DynExpr { } bool is_constant() const override { - return lhs->is_constant() && rhs->is_constant() && rhs->get_val() != 0 && - lhs->get_val() % rhs->get_val() == 0; + return lhs->is_constant() && rhs->is_constant() && rhs->get_val() != 0; } int64_t get_val() const override { - CHECK(is_constant()) << "Attempted to get integer value of non-integral " - << "division expression"; + CHECK(is_constant()) + << "Attempted to evaluate a non-constant or zero-divisor expression"; return lhs->get_val() / rhs->get_val(); } @@ -513,7 +528,9 @@ class Div : public DynExpr { if (lhs->is_dynamic() && rhs->is_dynamic()) return std::nullopt; if (lhs->get_all_ids().size() == 1 && rhs->is_constant()) { // (A / c) = x <=> A = x * c => solve A = y with y = x * c - return lhs->solve(x * rhs->get_val()); + const int64_t divisor = rhs->get_val(); + if (divisor == 0) return std::nullopt; + return lhs->solve(x * divisor); } if (rhs->get_all_ids().size() == 1 && lhs->is_constant()) { // (c / A) = x <=> A = c / x => solve A = y with y = c / x @@ -529,8 +546,166 @@ 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, int64_t k); DynExpr* operator/(DynExpr& lhs, DynExpr& rhs); DynExpr* operator/(DynExpr& lhs, int64_t d); DynExpr* operator+(DynExpr& lhs, DynExpr& rhs); @@ -546,6 +721,9 @@ inline DExpr operator*(const DExpr& lhs, const DExpr& rhs) { inline DExpr operator*(int64_t lhs, const DExpr& rhs) { return DExpr::Adopt(lhs * *rhs.get()); } +inline DExpr operator*(const DExpr& lhs, int64_t rhs) { + return DExpr::Adopt(*lhs.get() * rhs); +} inline DExpr operator/(const DExpr& lhs, const DExpr& rhs) { return DExpr::Adopt(*lhs.get() / *rhs.get()); } @@ -593,6 +771,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..e79fd805c1bfb5 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,115 @@ TEST_F(ShapeTest, DExprSimplifyCombinesEqualFractions) { EXPECT_EQ("A", DExprToString(expr.simplify())); } +TEST_F(ShapeTest, DExprSolveRejectsZeroDivisor) { + DExpr expr = DExpr::Var(1) / DExpr::Const(0); + EXPECT_FALSE(expr->solve(7).has_value()); +} + +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 positive_affine = + DExpr::Max(2 * DExpr::Var(1) - 1, DExpr::Const(0)); + EXPECT_EQ("((2 * A) - 1)", + DExprToString(positive_affine.simplify())); + + DExpr unbounded_below = + DExpr::Max(DExpr::Var(1) - DExpr::Var(2), DExpr::Const(0)); + EXPECT_EQ("max((A + ((-1) * B)), 0)", + DExprToString(unbounded_below.simplify())); + + DExpr evaluated = expr.substitute(1, DExpr::Const(7)).simplify(); + EXPECT_EQ(DExpr::Kind::kConstant, evaluated.kind()); + EXPECT_EQ(7, evaluated->get_val()); + + DExpr divided = + DExpr::Max((DExpr::Var(1) + 1) / 2, DExpr::Const(0)); + DExpr divided_evaluated = + divided.substitute(1, DExpr::Const(100)).simplify(); + EXPECT_EQ(DExpr::Kind::kConstant, divided_evaluated.kind()); + EXPECT_EQ(50, divided_evaluated->get_val()); + + ExpressionProto proto; + expr.to_proto(&proto); + EXPECT_TRUE(expr == DExprFromProto(proto)); +} + +TEST_F(ShapeTest, DExprSelectUsesDynamicPredicate) { + DExpr delta = DExpr::Var(1) - 3; + DExpr expr = DExpr::Select(DExpr::Gt(delta, DExpr::Const(0)), + DExpr::Const(7), DExpr::Const(11)); + EXPECT_EQ("select(((A - 3) > 0), 7, 11)", + DExprToString(expr.simplify())); + EXPECT_EQ(7, expr.substitute(1, DExpr::Const(5))->s()->get_val()); + EXPECT_EQ(11, expr.substitute(1, DExpr::Const(1))->s()->get_val()); + + ExpressionProto proto; + expr.to_proto(&proto); + EXPECT_TRUE(expr == DExprFromProto(proto)); +} + +TEST_F(ShapeTest, DExprUnknownPropagatesThroughGtAndSelect) { + DExpr gt = DExpr::Gt(DExpr::Unknown(), DExpr::Const(0)).simplify(); + EXPECT_TRUE(gt.is_unknown()); + + DExpr select = + DExpr::Select(DExpr::Unknown(), DExpr::Const(7), DExpr::Const(11)) + .simplify(); + EXPECT_TRUE(select.is_unknown()); +} + +TEST_F(ShapeTest, DExprConstantDivisionIsNotDynamic) { + const DExpr expr = DExpr::Const(7) / DExpr::Const(3); + + EXPECT_FALSE(expr->is_dynamic()); + EXPECT_EQ(expr->get_val(), 2); +} + +TEST_F(ShapeTest, DExprFindsSmallestSubexpressionCoveringAllVariables) { + const DExpr shared_core = DExpr::Var(1) + DExpr::Var(2); + const DExpr expr = (shared_core + 2) * (shared_core + 3); + + EXPECT_EQ( + shared_core, + expr.find_smallest_subexpression_covering_all_variables()); +} + +TEST_F(ShapeTest, DExprFindsCommonCoreAcrossRepeatedBranches) { + const DExpr shared_core = DExpr::Var(1) + DExpr::Var(2); + + EXPECT_EQ(shared_core, + (shared_core * shared_core) + .find_smallest_subexpression_covering_all_variables()); + EXPECT_EQ(shared_core, + (shared_core * (DExpr::Const(3) + shared_core)) + .find_smallest_subexpression_covering_all_variables()); +} + +TEST_F(ShapeTest, DExprUsesParentWhenChildCoresDiffer) { + const DExpr complete_expr = + (DExpr::Var(1) + DExpr::Var(2)) * + (DExpr::Var(1) - DExpr::Var(2)); + + EXPECT_EQ( + complete_expr, + complete_expr.find_smallest_subexpression_covering_all_variables()); +} + +TEST_F(ShapeTest, DExprReplacesEveryMatchingSubexpression) { + const DExpr shared_core = DExpr::Var(1) + DExpr::Var(2); + const DExpr expr = (shared_core + 2) * (shared_core + 3); + const DExpr replacement = DExpr::Var(7); + + EXPECT_EQ((replacement + 2) * (replacement + 3), + expr.replace_subexpression(shared_core, replacement)); +} + 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/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..5a01598ad7ae1c 100644 --- a/third_party/xla/xla/shape_util_test.cc +++ b/third_party/xla/xla/shape_util_test.cc @@ -1049,7 +1049,8 @@ TEST(ShapeUtilTest, InvalidDynamicDimension) { TEST(ShapeUtilTest, PermuteDynamicDimensions) { Shape shape = ShapeUtil::MakeShape(F32, {10, 100, 1000}, - /*dynamic_dimensions*/ {false, true, true}); + /*dynamic_dimensions*/ {false, true, true}, + /*expressions=*/{}); SCOPED_TRACE(absl::StrCat("shape=", shape.ToString())); std::vector permutation(3); @@ -1129,6 +1130,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 +1163,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}); 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; +}