You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将类鸢尾花分类的TensorFlow模型导出为.tflite文件?

把预构建Estimator模型转成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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.29 07:53:14