tf.Keras自定义损失与指标函数如何获取原始targets
可行的实现方案共3种,可根据你的业务场景选择:
方案1:把原始targets作为模型的额外输入
- 改动最小,不需要调整现有训练逻辑的核心流程。你可以在调用
train_on_batch时,把x参数修改为[input_batch, original_targets],相当于把原始标签作为模型的一个辅助输入传入。 - 定义自定义损失/指标的时候,直接从模型的输入张量里取出原始targets即可,不需要依赖
y_true传值。 - 注意模型定义的时候要声明两个输入:主输入和原始标签输入,原始标签输入不需要参与前向计算的主逻辑,仅需传给损失和指标函数使用即可。
方案2:改用自定义训练循环
- 完全抛弃内置的
train_on_batch、fit等接口的封装限制,自己编写每一步的训练逻辑:- 手动执行前向传播得到
y_pred - 直接同时把
original_targets、处理后的目标值、y_pred传给你的损失和指标函数 - 手动计算梯度、调用优化器更新参数
- 手动执行前向传播得到
- 该方案自由度最高,不需要修改模型输入结构,适合本身就需要自定义训练流程的场景。
方案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
相关产品推荐
相关产品推荐

