新手求助:如何在Google Colab下载TensorFlow训练模型为.pb文件
嘿,作为刚入门机器学习的新手,要把Colab里训练好的模型导出给Android做离线预测确实需要明确的步骤,我给你一步步拆解,保证你能顺利拿到需要的.pb文件:
先明确个小知识点
首先说下:Android端其实更常用TensorFlow Lite的.tflite格式(体积更小、运行更快),但你指定要.pb格式,我会同时覆盖包含权重的SavedModel格式(带saved_model.pb)和单一冻结.pb文件两种方式,你可以按需选择。
步骤1:确认你的模型训练完成
首先确保Colab里的模型已经训练完毕,比如你有一个叫model的Keras/TensorFlow模型变量,能正常做预测。
步骤2:导出SavedModel格式(自带saved_model.pb)
SavedModel是TensorFlow的标准保存格式,里面会自动生成saved_model.pb(模型计算图)和variables文件夹(模型权重),Android可以直接加载这个格式,操作超简单:
- 在Colab里运行这段代码(替换成你的模型变量):
import tensorflow as tf # 定义保存路径,Colab的临时目录就行 save_dir = '/content/my_trained_model' # 保存模型 tf.saved_model.save(model, save_dir) - 运行完后,点击Colab左侧的文件夹图标(文件管理器),就能找到
/content/my_trained_model文件夹,里面就有你要的saved_model.pb啦。
下载这个文件夹到本地
- 右键点击
my_trained_model文件夹,选择Download,Colab会自动把它打包成.zip文件下载到你的电脑 - 解压后就能得到
saved_model.pb和variables文件夹,这俩要一起放到Android项目里哦
可选:导出单一冻结.pb文件(权重+计算图合并)
如果你想要一个独立的.pb文件(把权重和计算图整合到一起),可以按下面的步骤来:
第一步:先按上面的方法导出SavedModel
第二步:生成冻结图
在Colab里运行这段代码:
import tensorflow as tf from tensorflow.python.framework.convert_to_constants import convert_variables_to_constants_v2 # 加载刚才保存的SavedModel loaded_model = tf.saved_model.load('/content/my_trained_model') # 获取模型的推理签名(默认是serving_default) infer_signature = loaded_model.signatures['serving_default'] # 转换成可冻结的ConcreteFunction concrete_func = infer_signature.get_concrete_function() # 把变量转换成常量,合并到计算图里 frozen_func = convert_variables_to_constants_v2(concrete_func) frozen_graph = frozen_func.graph # 保存冻结后的.pb文件 tf.io.write_graph( graph_or_graph_def=frozen_graph, logdir='/content/', name='frozen_model.pb', as_text=False )
运行完后,你会在Colab的/content/目录下看到frozen_model.pb,这个就是单一的、包含所有权重的.pb文件。
下载冻结.pb文件
直接在文件管理器里找到frozen_model.pb,右键选择Download就能直接下载到本地。
小验证:确保.pb文件能用(可选)
如果你想确认导出的.pb文件没问题,可以在本地跑个小测试(需要本地装TensorFlow):
import tensorflow as tf # 加载冻结的.pb文件 with tf.io.gfile.GFile('/你本地的路径/frozen_model.pb', 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) # 启动会话加载图 with tf.compat.v1.Session() as sess: sess.graph.as_default() tf.import_graph_def(graph_def, name='') # 替换成你模型的输入、输出节点名称 # 你可以在Colab里用model.input.name和model.output.name查看 input_tensor = sess.graph.get_tensor_by_name('input_1:0') output_tensor = sess.graph.get_tensor_by_name('dense_1/Softmax:0') # 用随机输入测试 test_input = tf.random.normal([1, 224, 224, 3]) # 替换成你模型的输入形状 output = sess.run(output_tensor, feed_dict={input_tensor: test_input.numpy()}) print("测试输出:", output)
给Android开发的小建议
如果之后你想优化Android端的性能,推荐把SavedModel转换成.tflite格式,代码也很简单:
converter = tf.lite.TFLiteConverter.from_saved_model('/content/my_trained_model') tflite_model = converter.convert() # 保存.tflite文件 with open('/content/my_model.tflite', 'wb') as f: f.write(tflite_model)
不管是.pb还是.tflite,放到Android项目的assets文件夹里就能加载使用啦。
如果你的Test.ipynb里有具体的模型代码(比如是CNN、MLP之类的),可以贴出来,我能帮你调整更适配的导出步骤~
内容的提问来源于stack exchange,提问作者peanut

