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

Siamese Network绘制ROC曲线及二分类场景代码适配问题咨询

孪生网络(Siamese Network)可以绘制ROC曲线

完全可以。孪生网络的核心任务是判断输入样本对是否匹配,属于典型的二分类任务,只要能获取模型输出的连续相似度/概率得分,以及对应的真实二分类标签,就能按照常规二分类任务的流程绘制ROC曲线,无需特殊适配。

混淆矩阵与ROC不匹配问题排查方案

你提供的ROC绘制代码本身逻辑是通用的,不需要专门针对二分类场景做结构修改,出现异常大概率是以下几个细节问题导致的:

  • 标签与模型输出的对应关系错误:Keras官方孪生网络示例中,模型输出的分值越高,代表两个输入样本差异越大、越不匹配,标签1对应「样本对不匹配」,标签0对应「样本对匹配」。如果你迁移到自己的二分类任务时,标签定义和这个逻辑相反,比如你定义标签1为「不变(Unchange)」,那用pred>0.5赋值为1的逻辑就完全颠倒,自然会出现结果不匹配的问题。
  • 未指定正样本类别:sklearn的roc_curve函数默认将标签中数值更大的类别作为正样本,如果你自定义的正样本标签不是数值更大的那一个,就会导致FPR和TPR计算完全颠倒,ROC曲线出现异常。需要在调用时手动指定pos_label参数为你定义的正样本的标签值。
  • 维度不匹配:你代码中pred = pred[:,0]是适配原示例模型输出为二维数组的情况,如果你自己修改后的模型输出本身就是一维数组,这一步会导致维度错误,你可以先打印pred.shape和labels_test.shape确认二者长度、维度完全一致。

调整后的参考代码

from sklearn.metrics import confusion_matrix,accuracy_score, roc_curve, auc
import seaborn as sns
import numpy as np
import matplotlib.pyplot as plt
sns.set_style("whitegrid")

pred = siamese.predict([x_test_1, x_test_2])
# 适配输出维度,确保是一维数组
if len(pred.shape) > 1:
    pred = pred[:,0]
# 阈值逻辑要和自己的标签定义匹配,如果得分越高越相似,pred>0.5应该对应匹配类(比如Unchange类)
pred_NN_01 = np.where(pred > 0.5, 1, 0) 
# 打印准确率
acc_NN = accuracy_score(labels_test, pred_NN_01)
print('Overall accuracy of Neural Network model:', acc_NN)

# 绘制ROC曲线,pos_label替换为自己任务中正样本的标签值,比如把Change作为正样本则pos_label=1
false_positive_rate, recall, thresholds = roc_curve(labels_test, pred, pos_label=1)
roc_auc = auc(false_positive_rate, recall)
plt.figure()
plt.title('Receiver Operating Characteristic (ROC)')
plt.plot(false_positive_rate, recall, 'b', label = 'AUC = %0.3f' %roc_auc)
plt.legend(loc='lower right')
plt.plot([0,1], [0,1], 'r--')
plt.xlim([0.0,1.0])
plt.ylim([0.0,1.0])
plt.ylabel('Recall')
plt.xlabel('Fall-out (1-Specificity)')
plt.show()

# 绘制混淆矩阵
cm = confusion_matrix(labels_test, pred_NN_01)
labels = ['Unchange', 'Change']
plt.figure(figsize=(8,6))
sns.heatmap(cm,xticklabels=labels, yticklabels=labels, annot=True, fmt='d', cmap="Blues", vmin = 0.2)
plt.title('Confusion Matrix')
plt.ylabel('True Class')
plt.xlabel('Predicted Class')
plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 05:09:01