Makes a TFLiteConverter object based on the flags provided. Args: flags: argparse.Namespace object containing TFLite flags. Returns: TFLiteConverter object. Raises: ValueError: Invalid flags.
(flags)
| 64 | |
| 65 | |
| 66 | def _get_toco_converter(flags): |
| 67 | """Makes a TFLiteConverter object based on the flags provided. |
| 68 | |
| 69 | Args: |
| 70 | flags: argparse.Namespace object containing TFLite flags. |
| 71 | |
| 72 | Returns: |
| 73 | TFLiteConverter object. |
| 74 | |
| 75 | Raises: |
| 76 | ValueError: Invalid flags. |
| 77 | """ |
| 78 | # Parse input and output arrays. |
| 79 | input_arrays = _parse_array(flags.input_arrays) |
| 80 | input_shapes = None |
| 81 | if flags.input_shapes: |
| 82 | input_shapes_list = [ |
| 83 | _parse_array(shape, type_fn=int) |
| 84 | for shape in flags.input_shapes.split(":") |
| 85 | ] |
| 86 | input_shapes = dict(zip(input_arrays, input_shapes_list)) |
| 87 | output_arrays = _parse_array(flags.output_arrays) |
| 88 | |
| 89 | converter_kwargs = { |
| 90 | "input_arrays": input_arrays, |
| 91 | "input_shapes": input_shapes, |
| 92 | "output_arrays": output_arrays |
| 93 | } |
| 94 | |
| 95 | # Create TFLiteConverter. |
| 96 | if flags.graph_def_file: |
| 97 | converter_fn = lite.TFLiteConverter.from_frozen_graph |
| 98 | converter_kwargs["graph_def_file"] = flags.graph_def_file |
| 99 | elif flags.saved_model_dir: |
| 100 | converter_fn = lite.TFLiteConverter.from_saved_model |
| 101 | converter_kwargs["saved_model_dir"] = flags.saved_model_dir |
| 102 | converter_kwargs["tag_set"] = _parse_set(flags.saved_model_tag_set) |
| 103 | converter_kwargs["signature_key"] = flags.saved_model_signature_key |
| 104 | elif flags.keras_model_file: |
| 105 | converter_fn = lite.TFLiteConverter.from_keras_model_file |
| 106 | converter_kwargs["model_file"] = flags.keras_model_file |
| 107 | else: |
| 108 | raise ValueError("--graph_def_file, --saved_model_dir, or " |
| 109 | "--keras_model_file must be specified.") |
| 110 | |
| 111 | return converter_fn(**converter_kwargs) |
| 112 | |
| 113 | |
| 114 | def _convert_tf1_model(flags): |
no test coverage detected