TensorFlow1.15 Java/Scala加载SavedModel时如何关闭learning_phase
TensorFlow 1.x 中Keras模型的learning_phase是图中一个固定名为keras_learning_phase:0的布尔型占位符,值为1/True代表启用训练模式,0/False代表禁用(推理模式)。加载模型后可以直接在会话中读取该节点值判断当前状态:
import org.tensorflow.Tensor var lpStatus: String = "未知" val lpTensor = try { session.runner() .fetch("keras_learning_phase:0") .run() .get(0) } catch { case e: IllegalArgumentException => lpStatus = "模型导出时已固定learning_phase,无分支切换逻辑" null } if (lpTensor != null) { val isTraining = lpTensor.booleanValue() lpStatus = if (isTraining) "已启用(训练模式)" else "已禁用(推理模式)" lpTensor.close() } println(s"当前learning_phase状态:$lpStatus")
如果执行时提示找不到keras_learning_phase:0节点,说明导出模型时已经固化了learning_phase值,不需要再做额外配置。
TensorFlow Java/Scala 1.15.0 版本完全支持该操作,和Python端tf.keras.backend.learning_phase(0)效果完全一致的实现方式有两种,按需选择即可。
方案1:推理时直接feed节点值(无需重导模型,推荐)
不需要修改现有模型导出逻辑,只需要在每次推理时,给keras_learning_phase:0占位符传入false值,即可强制模型走推理分支,和Python端行为完全对齐。修改后的预测代码如下:
val input1: Tensor[_] = Tensor.create(embed(name1)) val input2: Tensor[_] = Tensor.create(embed(name2)) // 构造learning_phase常量,false对应推理模式 val lpFlag: Tensor[_] = Tensor.create(false) val result: Tensor[_] = session.runner() .fetch("StatefulPartitionedCall:0") .feed("serving_default_input_1:0", input1) .feed("serving_default_input_2:0", input2) // 新增行:固定为推理模式 .feed("keras_learning_phase:0", lpFlag) .run() .get(0) // 处理result的取值逻辑 // ... // 所有张量用完必须手动close,避免堆外内存泄漏 Seq(input1, input2, lpFlag, result).foreach(_.close())
该方案是TF Java生态中处理Keras模型推理的标准做法,不存在API兼容问题,可以解决Dropout、BatchNorm层在训练/推理模式下计算逻辑不一致导致的结果偏差问题。
方案2:导出模型时固化learning_phase(一劳永逸)
如果可以重新走模型导出流程,在Python端导出SavedModel前直接设置全局learning_phase为0,导出的模型会直接剪掉训练分支,不存在keras_learning_phase:0占位符,Java/Scala侧加载后直接推理即可,不需要额外传值:
import tensorflow as tf # 导出模型前先执行这行,固定为推理模式 tf.keras.backend.set_learning_phase(0) # 后续正常执行模型导出逻辑即可 # model.save("/Models/lstm_model_fullname_with_watchlist")
注意:TensorFlow Java 1.x没有暴露Keras后端修改全局learning_phase默认值的接口,不要尝试通过反射或者其他方式修改图默认值,直接feed节点是最稳定可靠的方案。
内容的提问来源于stack exchange,提问作者Sudha Rajamanickam

