如何从缓存数据集获取Tab Transformer预测值以生成混淆矩阵与ROC_AUC图?
解决缓存数据集无法拆分以生成混淆矩阵和ROC_AUC图的问题
核心思路
TensorFlow的缓存数据集(tf.data.Dataset)无需拆分为单独的X和y,直接通过遍历数据集提取真实标签,同时用模型对整个数据集做预测即可。
具体实现步骤
提取真实标签与预测结果
不需要拆分数据集,直接遍历缓存集获取真实标签,用model.predict()直接接收数据集对象生成预测值:import numpy as np # 处理训练集 # 提取训练集真实标签 y_train_true = [] for _, y in train_dataset.unbatch(): y_train_true.append(y.numpy()) y_train_true = np.array(y_train_true) # 获取训练集预测概率并转换为二分类标签 y_predicted_train = model.predict(train_dataset) y_pred_train = (y_predicted_train > 0.5).astype(int).flatten() # 处理测试集 # 提取测试集真实标签 y_test_true = [] for _, y in test_dataset.unbatch(): y_test_true.append(y.numpy()) y_test_true = np.array(y_test_true) # 获取测试集预测概率并转换为二分类标签 y_predicted_test = model.predict(test_dataset) y_pred_test = (y_predicted_test > 0.5).astype(int).flatten()注:如果数据集规模较大,
unbatch()内存占用过高,可分批遍历后拼接结果生成混淆矩阵与ROC-AUC图
用sklearn的评估工具和matplotlib完成可视化:from sklearn.metrics import confusion_matrix, roc_curve, auc import matplotlib.pyplot as plt # 输出混淆矩阵 print("训练集混淆矩阵:") print(confusion_matrix(y_train_true, y_pred_train)) print("\n测试集混淆矩阵:") print(confusion_matrix(y_test_true, y_pred_test)) # 定义ROC-AUC绘图函数 def plot_roc_auc(y_true, y_pred_proba, title): fpr, tpr, _ = roc_curve(y_true, y_pred_proba) roc_auc = auc(fpr, tpr) plt.figure(figsize=(8, 6)) plt.plot(fpr, tpr, color='#ff7f0e', lw=2, label=f'ROC曲线 (AUC = {roc_auc:.2f})') plt.plot([0, 1], [0, 1], color='#1f77b4', lw=2, linestyle='--') plt.xlim([0.0, 1.0]) plt.ylim([0.0, 1.05]) plt.xlabel('假阳性率') plt.ylabel('真阳性率') plt.title(title) plt.legend(loc='lower right') plt.show() # 绘制训练集与测试集ROC-AUC图 plot_roc_auc(y_train_true, y_predicted_train.flatten(), '训练集ROC-AUC曲线') plot_roc_auc(y_test_true, y_predicted_test.flatten(), '测试集ROC-AUC曲线')关键注意点
model.predict()支持直接传入tf.data.Dataset,会自动处理批次,无需手动拆分输入特征- 确保数据集的结构(输入特征维度、标签格式)与模型输入要求匹配,避免预测时抛出维度不匹配错误
内容的提问来源于stack exchange,提问作者user19634316
相关产品推荐
相关产品推荐

