| 399 | } |
| 400 | |
| 401 | Status AddFlatMapNode(const string& input_dataset, |
| 402 | gtl::ArraySlice<string> other_arguments, |
| 403 | gtl::ArraySlice<DataType> t_arguments, |
| 404 | const FunctionDef& flat_map_fn, |
| 405 | const AttrValue& output_shapes, |
| 406 | const DataTypeVector& output_types, |
| 407 | FunctionLibraryDefinition* flib, MutableGraphView* graph, |
| 408 | NodeDef** result) { |
| 409 | TF_RETURN_IF_ERROR(flib->AddFunctionDef(flat_map_fn)); |
| 410 | AttrValue f; |
| 411 | f.mutable_func()->set_name(flat_map_fn.signature().name()); |
| 412 | |
| 413 | NodeDef flat_map_node; |
| 414 | flat_map_node.set_op("FlatMapDataset"); |
| 415 | flat_map_node.add_input(input_dataset); |
| 416 | for (const auto& arg : other_arguments) { |
| 417 | flat_map_node.add_input(arg); |
| 418 | } |
| 419 | AddNodeAttr("f", f, &flat_map_node); |
| 420 | AddNodeAttr("Targuments", t_arguments, &flat_map_node); |
| 421 | AddNodeAttr(kOutputShapesAttr, output_shapes, &flat_map_node); |
| 422 | AddNodeAttr(kOutputTypesAttr, output_types, &flat_map_node); |
| 423 | |
| 424 | graph_utils::SetUniqueGraphNodeName("rebatch/flat_map", graph->graph(), |
| 425 | &flat_map_node); |
| 426 | *result = graph->AddNode(std::move(flat_map_node)); |
| 427 | return Status::OK(); |
| 428 | } |
| 429 | |
| 430 | // def flat_map_fn(*batched_components): |
| 431 | // batch_size = tf.shape(batched_components[0])[0] |
no test coverage detected