put onnx to tflite conversion behind flag in training code

This commit is contained in:
david.scripka 2025-12-30 11:47:22 -05:00
parent af923e1d57
commit 368c03716d

View file

@ -630,6 +630,13 @@ if __name__ == '__main__':
default="False", default="False",
required=False required=False
) )
parser.add_argument(
"--convert_to_tflite",
help="Convert the trained ONNX model to TFLite format",
action="store_true",
default="False",
required=False
)
args = parser.parse_args() args = parser.parse_args()
config = yaml.load(open(args.training_config, 'r').read(), yaml.Loader) config = yaml.load(open(args.training_config, 'r').read(), yaml.Loader)
@ -898,5 +905,6 @@ if __name__ == '__main__':
oww.export_model(model=best_model, model_name=config["model_name"], output_dir=config["output_dir"]) oww.export_model(model=best_model, model_name=config["model_name"], output_dir=config["output_dir"])
# Convert the model from onnx to tflite format # Convert the model from onnx to tflite format
convert_onnx_to_tflite(os.path.join(config["output_dir"], config["model_name"] + ".onnx"), if args.convert_to_tflite:
os.path.join(config["output_dir"], config["model_name"] + ".tflite")) convert_onnx_to_tflite(os.path.join(config["output_dir"], config["model_name"] + ".onnx"),
os.path.join(config["output_dir"], config["model_name"] + ".tflite"))