如何将类鸢尾花分类的TensorFlow模型导出为.tflite文件?
好消息,针对你用的这种基于预构建DNNClassifier的模型(和鸢尾花分类示例结构一致),完全不用走“导出GraphDef→冻结模型”那套繁琐流程,直接用Estimator自带的导出工具生成SavedModel,再转TFLite就可以了,步骤简单很多,我给你详细拆解:
步骤1:修改训练代码,导出SavedModel
在你的训练代码里(对应鸢尾花示例的premade_estimator.py),训练和评估完成后,添加几行代码就能导出标准的SavedModel——这是TensorFlow官方推荐的模型归档格式,TFLite转换器可以直接识别它。
具体代码修改
在eval_result打印完成后,新增以下代码:
# 导出SavedModel到指定目录 export_dir_base = "./football_saved_model" # 根据你的特征列生成服务用的输入接收器 serving_input_receiver_fn = tf.estimator.export.build_parsing_serving_input_receiver_fn( tf.feature_column.make_parse_example_spec(my_feature_columns) ) # 执行导出 export_path = classifier.export_savedmodel(export_dir_base, serving_input_receiver_fn) print(f"SavedModel已导出到: {export_path}")
这段代码的作用是:
- 利用你的特征列自动生成符合TFLite要求的输入格式
- 将训练好的模型参数和计算图完整保存到
football_saved_model目录下,生成一个带时间戳数字的子目录(比如1699999999),这个就是我们要用来转换的模型目录。
步骤2:用tflite_convert转成.tflite文件
导出SavedModel后,直接用TensorFlow自带的命令行工具tflite_convert完成转换,不需要手动处理pb文件。
针对TensorFlow 1.8.0的命令
打开终端,运行:
tflite_convert \ --output_file=./football_action_model.tflite \ --saved_model_dir=./football_saved_model/[你的数字版本目录]
注意把[你的数字版本目录]替换成你导出后生成的那个带数字的子目录名(比如1699999999)。
(可选)指定输入输出节点
如果转换器自动检测输入输出有问题,可以先通过以下命令查看SavedModel的输入输出张量名:
saved_model_cli show --dir ./football_saved_model/[数字版本目录] --all
找到对应的输入和输出名称后,在转换命令里加上:
--input_arrays=inputs \ --output_arrays=dnn/head/predictions/probabilities
(这里的名字是鸢尾花模型的默认值,你的足球模型可能类似,具体以saved_model_cli的输出为准)
适配你的足球动作模型的注意点
因为你的模型和鸢尾花示例的差异,只需要确保这几点:
- 你的数据集加载函数(对应
iris_data.load_data())返回的train_x包含10个特征的键值对,特征列会自动循环生成,不用手动修改 DNNClassifier的n_classes=6参数正确设置,导出时会自动包含到模型里,转换TFLite时无需额外调整
完整修改后的核心代码片段
给你贴一下修改后的main函数核心部分,方便你参考:
def main(argv): args = parser.parse_args(argv[1:]) # 替换成你的足球动作数据集加载逻辑 (train_x, train_y), (test_x, test_y) = your_football_data.load_data() # 自动生成10个特征的特征列 my_feature_columns = [] for key in train_x.keys(): my_feature_columns.append(tf.feature_column.numeric_column(key=key)) # 构建你的DNN分类器(n_classes=6对应你的需求) classifier = tf.estimator.DNNClassifier( feature_columns=my_feature_columns, hidden_units=[10, 10], # 可根据你的需求调整隐藏层节点数 n_classes=6) # 训练模型 classifier.train( input_fn=lambda: your_football_data.train_input_fn(train_x, train_y, args.batch_size), steps=args.train_steps) # 评估模型 eval_result = classifier.evaluate( input_fn=lambda: your_football_data.eval_input_fn(test_x, test_y, args.batch_size)) print('\nTest set accuracy: {accuracy:0.3f}\n'.format(**eval_result)) # --- 新增:导出SavedModel --- export_dir_base = "./football_saved_model" serving_input_receiver_fn = tf.estimator.export.build_parsing_serving_input_receiver_fn( tf.feature_column.make_parse_example_spec(my_feature_columns) ) export_path = classifier.export_savedmodel(export_dir_base, serving_input_receiver_fn) print('SavedModel exported to:', export_path) # 原有的预测代码可以保留,不影响模型导出 # ...(这里放你的预测逻辑)
为什么这个方法更高效?
因为Estimator作为TensorFlow的高级API,已经封装了模型保存的所有细节——它会自动处理变量固化、计算图整理,生成的SavedModel是完整可迁移的,TFLite转换器可以直接读取,省去了手动导出GraphDef、调用freeze_graph工具这些容易出错的步骤,这些步骤其实是给用低级API(比如直接构建tf.Graph)写的模型用的,你的场景完全不需要。
内容的提问来源于stack exchange,提问作者Alexandru Balotescu

