如何转换TensorFlow1.15 TPU训练模型实现GPU部署加载
你遇到的两个报错原因非常明确:
- 第一次加载失败是因为你手里只有TPU训练产出的checkpoint三件套,从来没有导出过SavedModel格式文件,自然会报文件不存在的错误
- 第二次加载失败是因为checkpoint配套的.meta文件里固化了TPU专属的分布式算子(比如你遇到的
TPUReplicatedInput,还有配套的TPUReplicatedOutput等),纯GPU/CPU环境没有这些算子的内核实现,直接导入计算图必然报错。
可行解决方法
方法一:有TPU环境时优先用官方路径导出(最稳定)
这是出错概率最低的方案,步骤如下:
- 准备和训练时版本完全一致的TF1.15环境,正常连接TPU运行时
- 不要直接导入训练时生成的.meta计算图,重新搭建单卡推理用的前向计算图,全部去掉TPU分布式策略、TPU算子相关的代码,只保留从输入到模型输出的纯计算逻辑,不需要反向传播相关节点
- 搭建完推理图后,从现有checkpoint加载权重,再导出为标准SavedModel格式,参考代码如下:
import tensorflow as tf # 替换成你自己的BioALBERT单卡模型构建逻辑,结构和训练时完全一致,仅删除TPU相关包装 def build_inference_graph(max_seq_len=128): input_ids = tf.placeholder(tf.int32, shape=[None, max_seq_len], name="input_ids") attention_mask = tf.placeholder(tf.int32, shape=[None, max_seq_len], name="attention_mask") # 调用ALBERT主干结构,输出你需要的预测结果/词向量 model_output = bioalbert_forward(input_ids, attention_mask, is_training=False) return input_ids, attention_mask, model_output tf.reset_default_graph() input_ids, attention_mask, model_output = build_inference_graph() saver = tf.train.Saver(var_list=tf.global_variables()) with tf.Session() as sess: # 加载checkpoint权重,注意路径只写前缀,不要加.data/.index后缀 saver.restore(sess, "./model.ckpt-best") # 导出通用SavedModel tf.saved_model.simple_save( sess, export_dir="./gpu_runnable_model", inputs={"input_ids": input_ids, "attention_mask": attention_mask}, outputs={"output": model_output} )
- 导出完成后,把生成的
gpu_runnable_model目录拷贝到纯GPU环境,就可以用常规的SavedModel加载接口正常推理,不会再出现TPU算子相关报错。
方法二:无TPU环境时手动适配权重
如果没法接入TPU运行环境,可以按以下步骤处理:
- 同样不要尝试直接导入.meta文件,先在GPU环境下重新搭建和训练时结构完全一致的单卡BioALBERT模型,删掉所有TPU相关的分布式包装代码(比如TPUEstimator、tpu.replicate相关调用)
- 图搭建完成后,用
tf.train.list_variables("./model.ckpt-best")打印checkpoint里存储的所有权重名称、张量形状,和你当前单卡模型里的变量做匹配,把权重名里带TPU副本标识的前缀(比如replica_0/、tpu_0/这类)替换成单卡图对应的变量名,做一层映射 - 用映射后的变量列表初始化Saver,加载checkpoint权重,加载完成后可以直接在当前会话推理,也可以重新导出为不带TPU算子的SavedModel格式留作后续部署用。
注意避坑
- 给
saver.restore传路径时,只需要传checkpoint的前缀名(也就是model.ckpt-best),不需要加.data-00000-of-00001这类后缀,TF会自动匹配配套的索引、数据文件 - 如果之前用TPUEstimator接口训练,导出推理模型时记得开 serving_only 模式,指定推理用的输入函数,导出的图会自动剥离训练阶段的TPU分布式算子,不需要手动改结构
- 如果加载权重时出现shape不匹配、变量找不到的问题,直接对比checkpoint里的变量列表和你当前图里的变量列表,调整权重名映射规则即可,不需要改模型结构。
内容的提问来源于stack exchange,提问作者Dima Lituiev
相关产品推荐
相关产品推荐

