| 161 | } // namespace |
| 162 | |
| 163 | TEST(FlexImportTest, ConditionalConst) { |
| 164 | Model model; |
| 165 | auto build_and_import_node = |
| 166 | [&model](const string& name, std::initializer_list<int64_t> shape, |
| 167 | tensorflow::DataType dtype, int64_t num_elements) { |
| 168 | NodeDef node; |
| 169 | BuildConstNode(shape, dtype, num_elements, &node); |
| 170 | node.set_name(name); |
| 171 | |
| 172 | const auto converter = internal::GetTensorFlowNodeConverterMapForFlex(); |
| 173 | return internal::ImportTensorFlowNode(node, TensorFlowImportFlags(), |
| 174 | ModelFlags(), &model, converter); |
| 175 | }; |
| 176 | |
| 177 | EXPECT_TRUE(build_and_import_node("Known", {1, 2, 3}, DT_INT32, 6).ok()); |
| 178 | EXPECT_TRUE(build_and_import_node("BadType", {1, 2, 3}, DT_INVALID, 6).ok()); |
| 179 | EXPECT_TRUE(build_and_import_node("Unknown", {1, -2, 3}, DT_INT32, 6).ok()); |
| 180 | |
| 181 | // We expect the "Known" node to be converted into an array, while the |
| 182 | // "Unknown" and "BadType" nodes are kept as operators. |
| 183 | EXPECT_EQ(model.operators.size(), 2); |
| 184 | EXPECT_TRUE(model.HasArray("Known")); |
| 185 | EXPECT_FALSE(model.HasArray("Unknown")); |
| 186 | EXPECT_FALSE(model.HasArray("BadType")); |
| 187 | } |
| 188 | |
| 189 | class ShapeImportTest : public ::testing::TestWithParam<tensorflow::DataType> { |
| 190 | }; |
nothing calls this directly
no test coverage detected