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

tf.Keras自定义损失与指标函数如何获取原始targets

可行的实现方案共3种,可根据你的业务场景选择:

方案1:把原始targets作为模型的额外输入

  • 改动最小,不需要调整现有训练逻辑的核心流程。你可以在调用train_on_batch时,把x参数修改为[input_batch, original_targets],相当于把原始标签作为模型的一个辅助输入传入。
  • 定义自定义损失/指标的时候,直接从模型的输入张量里取出原始targets即可,不需要依赖y_true传值。
  • 注意模型定义的时候要声明两个输入:主输入和原始标签输入,原始标签输入不需要参与前向计算的主逻辑,仅需传给损失和指标函数使用即可。

方案2:改用自定义训练循环

  • 完全抛弃内置的train_on_batch、fit等接口的封装限制,自己编写每一步的训练逻辑:
    1. 手动执行前向传播得到y_pred
    2. 直接同时把original_targets、处理后的目标值、y_pred传给你的损失和指标函数
    3. 手动计算梯度、调用优化器更新参数
  • 该方案自由度最高,不需要修改模型输入结构,适合本身就需要自定义训练流程的场景。

方案3:用闭包封装损失/指标函数,动态传入原始targets

  • 定义损失函数的时候用外层函数持有原始targets的变量,每步训练前更新这个变量的值即可,示例代码如下:
import tensorflow as tf

def loss_wrapper():
    # 静态图模式下用tf.Variable存储可避免计算图追踪问题
    original_targets = tf.Variable(tf.zeros(shape=(), dtype=tf.float32), shape=None)
    def custom_loss(y_true, y_pred):
        # 此处可同时使用y_true(处理后的目标值)、original_targets(原始标签)、y_pred计算损失
        return ... 
    def update_targets(new_targets):
        original_targets.assign(new_targets)
    return custom_loss, update_targets

# 初始化阶段拿到损失函数和目标值更新方法
custom_loss, update_targets = loss_wrapper()
# 每步训练前先调用update_targets传入当前步的原始targets
update_targets(original_targets)
# 再正常调用train_on_batch即可
model.train_on_batch(x=input_batch, y=someFunction(targets))
  • 该方案不需要修改模型结构也不需要更换训练逻辑,仅需调整损失/指标的定义方式即可。

内容的提问来源于stack exchange,提问作者Sid Anand

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 06:36:03