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")