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

Keras自定义ROC曲线指标报错:K.eval(y_true)触发InvalidArgumentError

Fixing ROC Calculation in Keras: Resolving InvalidArgumentError & Epoch-End Alternative

Hey there! Let's break down why you're seeing that InvalidArgumentError and walk through two solid solutions to get your ROC metrics working in Keras.

Why Your Current Code Fails

The error pops up because K.eval(y_true) tries to pull actual values from symbolic tensors during the model compilation phase. At this point, y_true and y_pred are just placeholder tensors—they don't have real data yet (that only flows in during training/evaluation). On top of that, Keras expects custom metrics to return a single scalar value, but roc_curve returns three arrays (fpr, tpr, thresholds), which wouldn't work as a training metric anyway.


Solution 1: Fix the Custom Metric (Symbolic ROC AUC)

If you want to track ROC AUC as a metric during training, you need to use symbolic operations (no eval() calls) that work with Keras/TensorFlow's graph. Here's a corrected version that calculates ROC AUC (the standard scalar metric for ROC performance):

import tensorflow as tf
from keras import backend as K

def compute_roc_auc(y_true, y_pred):
    # Use TensorFlow's symbolic AUC calculation
    # tf.metrics.auc returns (current_auc, update_op); we need the update_op to get the correct value
    auc = tf.metrics.auc(y_true, y_pred, num_thresholds=200, pos_label=2)[1]
    # Initialize local variables required by TensorFlow's metric
    K.get_session().run(tf.local_variables_initializer())
    return auc

You can use this directly in model.compile() like so:

sgd = SGD(lr=0.001, decay=1e-6, momentum=0.9, nesterov=True)
disc.compile(loss='categorical_crossentropy', optimizer=sgd, metrics=['accuracy', compute_roc_auc])

Solution 2: Calculate ROC After Each Epoch (Callback Approach)

If you need the full ROC curve (fpr, tpr, thresholds) instead of just AUC, the best approach is to use a Keras Callback to compute it after each epoch. This lets you use scikit-learn's roc_curve with real validation data:

from sklearn.metrics import roc_curve, auc
from keras.callbacks import Callback

class ROCEpochCallback(Callback):
    def __init__(self, val_data):
        # Store validation data (x_val, y_val)
        self.x_val, self.y_val = val_data

    def on_epoch_end(self, epoch, logs=None):
        # Get model predictions for validation set
        y_pred = self.model.predict(self.x_val, verbose=0)
        # Calculate full ROC curve and AUC
        fpr, tpr, thresholds = roc_curve(self.y_val.ravel(), y_pred.ravel(), pos_label=2)
        roc_auc = auc(fpr, tpr)
        
        # Print results (or save to file/logs)
        print(f"\nEpoch {epoch + 1} - Validation ROC AUC: {roc_auc:.4f}")
        # Optional: Save ROC curve data to a file
        # with open(f"roc_epoch_{epoch+1}.txt", "w") as f:
        #     for fp, tp in zip(fpr, tpr):
        #         f.write(f"{fp}\t{tp}\n")

# Usage example:
# Assume you have x_val, y_val as your validation dataset
roc_callback = ROCEpochCallback(val_data=(x_val, y_val))
disc.fit(x_train, y_train, epochs=10, validation_data=(x_val, y_val), callbacks=[roc_callback])

This approach is perfect if you want to visualize the ROC curve later or need detailed threshold data—since it uses real, post-epoch predictions instead of symbolic tensors.


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 10:15:23