Import a SavedModel into a TF 1.x-style graph and run `signature_key`.
(
save_dir, inputs,
signature_key=signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY)
| 71 | |
| 72 | |
| 73 | def _import_and_infer( |
| 74 | save_dir, inputs, |
| 75 | signature_key=signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY): |
| 76 | """Import a SavedModel into a TF 1.x-style graph and run `signature_key`.""" |
| 77 | graph = ops.Graph() |
| 78 | with graph.as_default(), session_lib.Session() as session: |
| 79 | model = loader.load(session, [tag_constants.SERVING], save_dir) |
| 80 | signature = model.signature_def[signature_key] |
| 81 | assert set(inputs.keys()) == set(signature.inputs.keys()) |
| 82 | feed_dict = {} |
| 83 | for arg_name in inputs.keys(): |
| 84 | feed_dict[graph.get_tensor_by_name(signature.inputs[arg_name].name)] = ( |
| 85 | inputs[arg_name]) |
| 86 | output_dict = {} |
| 87 | for output_name, output_tensor_info in signature.outputs.items(): |
| 88 | output_dict[output_name] = graph.get_tensor_by_name( |
| 89 | output_tensor_info.name) |
| 90 | return session.run(output_dict, feed_dict=feed_dict) |
| 91 | |
| 92 | |
| 93 | class SaveTest(test.TestCase): |
no test coverage detected