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

TensorFlow1.15 Java/Scala加载SavedModel时如何关闭learning_phase

检查模型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值,不需要再做额外配置。

Scala/Java 环境禁用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 00:51:23