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

如何在Keras自定义指标中正确使用训练数据input_prices?

解决Keras函数式模型中自定义指标使用额外输入的问题

你遇到的错误核心是:Keras函数式模型中的符号张量(KerasTensor)不能直接传入原生TensorFlow函数,必须通过Keras提供的Layer或Metric类封装处理。以下是两种可行的实现方案:


方案一:继承tf.keras.metrics.Metric实现可追踪状态的自定义指标

适合需要跨批次累加计算指标的场景,符合Keras官方规范:

import tensorflow as tf
from tensorflow.keras import layers, Model, metrics, losses

class OverpriceMetric(tf.keras.metrics.Metric):
    def __init__(self, name='overprice', **kwargs):
        super().__init__(name=name, **kwargs)
        # 定义用于累加的权重变量
        self.total_overprice = self.add_weight(name='total_overprice', initializer='zeros')
        self.sample_count = self.add_weight(name='sample_count', initializer='zeros')

    def update_state(self, y_true, y_pred, input_prices=None, sample_weight=None):
        if input_prices is None:
            raise ValueError("计算指标必须传入input_prices参数")
        
        # 计算单批次的超价差值
        y_pred_rounded = tf.round(y_pred)
        pred_total = tf.reduce_sum(y_pred_rounded * input_prices, axis=1)
        true_total = tf.reduce_sum(y_true * input_prices, axis=1)
        batch_overprice = pred_total - true_total
        
        # 累加总差值和样本数量
        self.total_overprice.assign_add(tf.reduce_sum(batch_overprice))
        self.sample_count.assign_add(tf.cast(tf.shape(batch_overprice)[0], tf.float32))

    def result(self):
        # 返回平均超价
        return self.total_overprice / self.sample_count

    def reset_state(self):
        # 重置指标状态(每个epoch或验证阶段开始时自动调用)
        self.total_overprice.assign(0.0)
        self.sample_count.assign(0.0)

模型编译与使用

由于该指标需要额外传入input_prices,需自定义训练步骤:

class KnapsackModel(Model):
    def __init__(self, item_count=5):
        super().__init__()
        self.item_count = item_count
        self.concat_layer = layers.Concatenate()
        self.dense_layer = layers.Dense(item_count, use_bias=False, activation="sigmoid")
        self.overprice_metric = OverpriceMetric()
        self.binary_acc = metrics.BinaryAccuracy()

    def call(self, inputs):
        input_weights, input_prices, input_capacity = inputs
        inputs_concat = self.concat_layer([input_weights, input_prices, input_capacity])
        return self.dense_layer(inputs_concat)

    def train_step(self, data):
        x, y = data
        input_weights, input_prices, input_capacity = x

        with tf.GradientTape() as tape:
            y_pred = self([input_weights, input_prices, input_capacity], training=True)
            loss = losses.binary_crossentropy(y, y_pred)

        # 更新梯度
        grads = tape.gradient(loss, self.trainable_variables)
        self.optimizer.apply_gradients(zip(grads, self.trainable_variables))

        # 更新指标
        self.binary_acc.update_state(y, y_pred)
        self.overprice_metric.update_state(y, y_pred, input_prices=input_prices)

        # 返回指标结果
        return {
            "loss": tf.reduce_mean(loss),
            "binary_accuracy": self.binary_acc.result(),
            "overprice": self.overprice_metric.result()
        }

# 初始化模型并编译
model = KnapsackModel(item_count=5)
model.compile(optimizer="sgd")

方案二:使用model.add_metric直接计算指标

适合无需跨批次累加,仅需输出单批次指标均值的场景,实现更简洁:

import tensorflow as tf
from tensorflow.keras import layers, Model, metrics, losses

def supervised_continues_knapsack(item_count=5):
    input_weights = layers.Input((item_count,))
    input_prices = layers.Input((item_count,))
    input_capacity = layers.Input((1,))
    
    inputs_concat = layers.Concatenate()([input_weights, input_prices, input_capacity])
    picks = layers.Dense(item_count, use_bias=False, activation="sigmoid")(inputs_concat)
    
    # 定义指标计算逻辑
    def calculate_overprice(y_true, y_pred):
        y_pred_rounded = tf.round(y_pred)
        pred_total = tf.reduce_sum(y_pred_rounded * input_prices, axis=1)
        true_total = tf.reduce_sum(y_true * input_prices, axis=1)
        return tf.reduce_mean(pred_total - true_total)
    
    # 创建模型并添加自定义指标
    model = Model(inputs=[input_weights, input_prices, input_capacity], outputs=[picks])
    # 绑定指标到模型输出,指定聚合方式为均值
    model.add_metric(calculate_overprice(model.output, model.output), name='overprice', aggregation='mean')
    
    # 编译模型,标准指标可直接放入metrics列表
    model.compile(optimizer="sgd",
                  loss=losses.binary_crossentropy,
                  metrics=[metrics.binary_accuracy])
    return model

使用说明

训练时,自定义指标overprice会自动出现在训练日志中,无需额外处理。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 03:15:03