Several improvements. ReplicateIdenticalInputs: * let it accept different input types, so long as they're in a subtyping relation, and infer the most specific one (rather than using the first non-empty argument it finds) * rename it to Merge UnaryContainerCreate: fix argument name forward_type_test.cc: fix mispelled name PiperOrigin-RevId: 418183008 Change-Id: I973d63e32227cd2ca98c8e69bbfc23d823195960
diff --git a/tensorflow/core/common_runtime/forward_type_inference_test.cc b/tensorflow/core/common_runtime/forward_type_inference_test.cc index a7c3bc7..45e3426 100644 --- a/tensorflow/core/common_runtime/forward_type_inference_test.cc +++ b/tensorflow/core/common_runtime/forward_type_inference_test.cc
@@ -386,10 +386,10 @@ // at least once after both its inputs have been resolved, so the graph always // has complete type information. EXPECT_THAT(status.error_message(), - ::testing::HasSubstr("expected identical input types")); + ::testing::HasSubstr("expected compatible input types")); } -TEST(WeakForwardTypeInferenceTest, ALwaysSucceeds) { +TEST(WeakForwardTypeInferenceTest, AlwaysSucceeds) { std::unique_ptr<Graph> graph(new Graph(OpRegistry::Global())); Scope root = Scope::NewRootScope().ExitOnError();
diff --git a/tensorflow/core/framework/full_type_inference_util.cc b/tensorflow/core/framework/full_type_inference_util.cc index ca1c1bd..f15218b 100644 --- a/tensorflow/core/framework/full_type_inference_util.cc +++ b/tensorflow/core/framework/full_type_inference_util.cc
@@ -41,15 +41,12 @@ }; } -// TODO(mdan): Rename to MergeIdenticalInputs. -ForwardTypeInferenceFn ReplicateIdenticalInputs() { +ForwardTypeInferenceFn Merge() { return [](const std::vector<std::reference_wrapper<const FullTypeDef>>& input_types) -> StatusOr<FullTypeDef> { DCHECK(!input_types.empty()); - FullTypeDef ret_type; - int first_known = -1; - FullTypeDef const* first_known_t = nullptr; + FullTypeDef merged; for (int i = 0; i < input_types.size(); i++) { const auto& t = input_types[i].get(); @@ -57,41 +54,42 @@ continue; } - if (first_known < 0) { - first_known = i; - first_known_t = &t; - *(ret_type.add_args()) = t; + if (IsSubtype(t, merged)) { + merged = t; + continue; + } + if (IsSubtype(merged, t)) { continue; } - // TODO(mdan): Make a deep comparison. - if (first_known_t->type_id() != t.type_id()) { - return Status( - error::INVALID_ARGUMENT, - absl::StrCat("expected identical input types, but input ", i, - " differed from ", first_known, ":\n", t.DebugString(), - "\nvs.\n", first_known_t->DebugString())); - } + return Status(error::INVALID_ARGUMENT, + absl::StrCat("expected compatible input types, but input ", + i, ":\n", t.DebugString(), + " is neither a subtype nor a supertype of the " + "combined inputs preceding it:\n", + merged.DebugString())); } - if (first_known >= 0) { + FullTypeDef ret_type; + if (merged.type_id() != TFT_UNSET) { ret_type.set_type_id(TFT_PRODUCT); + *(ret_type.add_args()) = merged; } return ret_type; }; } -ForwardTypeInferenceFn UnaryContainerCreate(FullTypeId t, int container_idx) { - return [t, container_idx]( +ForwardTypeInferenceFn UnaryContainerCreate(FullTypeId t, int element_idx) { + return [t, element_idx]( const std::vector<std::reference_wrapper<const FullTypeDef>>& input_types) -> StatusOr<FullTypeDef> { - DCHECK(input_types.size() >= container_idx); + DCHECK(input_types.size() >= element_idx); FullTypeDef ret_type; ret_type.set_type_id(TFT_PRODUCT); FullTypeDef* arg_t = ret_type.add_args(); arg_t->set_type_id(t); - *(arg_t->add_args()) = input_types[container_idx].get(); + *(arg_t->add_args()) = input_types[element_idx].get(); return ret_type; };
diff --git a/tensorflow/core/framework/full_type_inference_util.h b/tensorflow/core/framework/full_type_inference_util.h index fbce510..94a863a 100644 --- a/tensorflow/core/framework/full_type_inference_util.h +++ b/tensorflow/core/framework/full_type_inference_util.h
@@ -47,19 +47,22 @@ // input. // The n arg allows multiple outputs, e.g. (T -> Product[T, T]). // TODO(mdan): Drop defaults for readability if more non-(0, 1) cases appear. +// TODO(mdan): Rename to just Replicate. ForwardTypeInferenceFn ReplicateInput(int i = 0, int n = 1); // Helper for a type inference function which has the same type as a variadic // number of inputs, e.g. (T, T -> Product[T]), (T, T, T -> Product[T]), etc. -// Assumes all inputs are of identical type. -ForwardTypeInferenceFn ReplicateIdenticalInputs(); +// Infers the meet of the input types, in the sense of type meets (see +// https://en.wikipedia.org/wiki/Join_and_meet). This implementation is +// simplified to require the two inputs are a subtype of another. +ForwardTypeInferenceFn Merge(); // Helper for the type inference counterpart of Unary, that is (U -> // PRODUCT[<t>[U]]), where <t> is parameterized by this factory, and U is the -// type of the input specified by container_idx. +// type of the input specified by element_idx. // Note: when we migrate to a more formal type definition of an op, these two // functions will naturally merge. -ForwardTypeInferenceFn UnaryContainerCreate(FullTypeId t, int container_idx); +ForwardTypeInferenceFn UnaryContainerCreate(FullTypeId t, int element_idx); // Helper for ops with semantics of adding an element to a container (<t>[T]), // that is (<t>[U], V -> PRODUCT[<t>[Union[U, V]]]), where <t> is parameterized
diff --git a/tensorflow/core/framework/full_type_inference_util_test.cc b/tensorflow/core/framework/full_type_inference_util_test.cc index dc64a14..b4bae38 100644 --- a/tensorflow/core/framework/full_type_inference_util_test.cc +++ b/tensorflow/core/framework/full_type_inference_util_test.cc
@@ -93,11 +93,11 @@ EXPECT_EQ(rt.type_id(), TFT_UNSET); } -TEST(ReplicateIdenticalInputs, Single) { +TEST(Merge, Single) { FullTypeDef t; t.set_type_id(TFT_ARRAY); - const auto ret = ReplicateIdenticalInputs()({t}); + const auto ret = Merge()({t}); TF_EXPECT_OK(ret.status()); const FullTypeDef& rt = ret.ValueOrDie(); @@ -106,11 +106,11 @@ EXPECT_EQ(rt.args(0).type_id(), TFT_ARRAY); } -TEST(ReplicateIdenticalInputs, Double) { +TEST(Merge, Double) { FullTypeDef t; t.set_type_id(TFT_ARRAY); - const auto ret = ReplicateIdenticalInputs()({t, t}); + const auto ret = Merge()({t, t}); TF_EXPECT_OK(ret.status()); const FullTypeDef& rt = ret.ValueOrDie(); @@ -119,65 +119,70 @@ EXPECT_EQ(rt.args(0).type_id(), TFT_ARRAY); } -TEST(ReplicateIdenticalInputs, Unset) { +TEST(Merge, Unset) { FullTypeDef t; t.set_type_id(TFT_UNSET); - const auto ret = ReplicateIdenticalInputs()({t}); + const auto ret = Merge()({t}); TF_EXPECT_OK(ret.status()); const FullTypeDef& rt = ret.ValueOrDie(); EXPECT_EQ(rt.type_id(), TFT_UNSET); } -TEST(ReplicateIdenticalInputs, UnsetComponents) { +TEST(Merge, UnsetComponents) { FullTypeDef t1; FullTypeDef t2; - const auto ret = ReplicateIdenticalInputs()({t1, t2}); + const auto ret = Merge()({t1, t2}); TF_EXPECT_OK(ret.status()); const FullTypeDef& rt = ret.ValueOrDie(); EXPECT_EQ(rt.type_id(), TFT_UNSET); } -TEST(ReplicateIdenticalInputs, UsesPartialInfo_FirstUnknown) { - FullTypeDef t1; - FullTypeDef t2; - t2.set_type_id(TFT_ARRAY); - - const auto ret = ReplicateIdenticalInputs()({t1, t2}); +void ExpectInferredArrayOfTensor(StatusOr<FullTypeDef> ret) { TF_EXPECT_OK(ret.status()); const FullTypeDef& rt = ret.ValueOrDie(); EXPECT_EQ(rt.type_id(), TFT_PRODUCT); ASSERT_EQ(rt.args_size(), 1); EXPECT_EQ(rt.args(0).type_id(), TFT_ARRAY); + ASSERT_EQ(rt.args(0).args_size(), 1); + EXPECT_EQ(rt.args(0).args(0).type_id(), TFT_TENSOR); } -TEST(ReplicateIdenticalInputs, UsesPartialInfo_SecondUnknown) { - FullTypeDef t1; - t1.set_type_id(TFT_ARRAY); - FullTypeDef t2; - - const auto ret = ReplicateIdenticalInputs()({t1, t2}); - TF_EXPECT_OK(ret.status()); - - const FullTypeDef& rt = ret.ValueOrDie(); - EXPECT_EQ(rt.type_id(), TFT_PRODUCT); - ASSERT_EQ(rt.args_size(), 1); - EXPECT_EQ(rt.args(0).type_id(), TFT_ARRAY); -} - -TEST(ReplicateIdenticalInputs, RejectsMismatched) { +TEST(Merge, RejectsMismatched) { FullTypeDef t1; t1.set_type_id(TFT_ARRAY); FullTypeDef t2; t2.set_type_id(TFT_TENSOR); - const auto ret = ReplicateIdenticalInputs()({t1, t2}); + const auto ret = Merge()({t1, t2}); EXPECT_THAT(ret.status().error_message(), - ::testing::HasSubstr("expected identical input types")); + ::testing::HasSubstr("expected compatible input types")); +} + +TEST(Merge, UsesPartialInfo) { + FullTypeDef t1; + FullTypeDef t2; + t2.set_type_id(TFT_ARRAY); + t2.add_args()->set_type_id(TFT_TENSOR); + + ExpectInferredArrayOfTensor(Merge()({t1, t2})); + ExpectInferredArrayOfTensor(Merge()({t2, t1})); +} + +TEST(Merge, SelectsMostSpecificOfSubtypes) { + FullTypeDef t1; + t1.set_type_id(TFT_ARRAY); + t1.add_args()->set_type_id(TFT_ANY); + FullTypeDef t2; + t2.set_type_id(TFT_ARRAY); + t2.add_args()->set_type_id(TFT_TENSOR); + + ExpectInferredArrayOfTensor(Merge()({t1, t2})); + ExpectInferredArrayOfTensor(Merge()({t2, t1})); } TEST(UnaryContainerCreate, Basic) {
diff --git a/tensorflow/core/ops/control_flow_ops.cc b/tensorflow/core/ops/control_flow_ops.cc index a84e651..9455fc8 100644 --- a/tensorflow/core/ops/control_flow_ops.cc +++ b/tensorflow/core/ops/control_flow_ops.cc
@@ -151,7 +151,7 @@ .Output("value_index: int32") .Attr("T: type") .Attr("N: int >= 1") - .SetForwardTypeFn(full_type::ReplicateIdenticalInputs()) + .SetForwardTypeFn(full_type::Merge()) .SetShapeFn(MergeShape); REGISTER_OP("RefMerge")