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

如何为该训练模型计算混淆矩阵?(附训练代码)

给你的Seq2Seq模型计算混淆矩阵的方法

嘿,针对你这个用Encoder-Decoder结构训练的Seq2Seq模型,计算混淆矩阵其实不难,只是因为是序列预测任务,得稍微处理下数据格式。下面一步步来操作:

1. 先获取模型的预测结果

首先你得用训练好的模型对目标数据集(比如训练时拆分出的验证集)做预测。因为你训练时用了validation_split=0.2,要么手动把原始数据拆分成训练/验证集,要么直接用模型对整个输入数据预测后再拆分验证部分。代码示例如下:

# 假设你已经准备好验证集输入:val_encoder_input, val_decoder_input
predictions = model.predict([val_encoder_input, val_decoder_input])

这里的predictions形状是(样本数, 序列长度, 类别数),和你的decoder_target_data结构完全对应。

2. 把独热编码转成类别索引

因为你用了categorical_crossentropy损失,目标数据decoder_target_data是独热编码格式的,得转成具体的类别索引才能和预测结果对比:

import numpy as np

# 真实标签:从独热编码转成类别索引
true_labels = np.argmax(val_decoder_target, axis=-1)
# 预测标签:取每个时间步概率最大的类别作为预测结果
pred_labels = np.argmax(predictions, axis=-1)

现在这两个变量的形状都是(样本数, 序列长度),每个位置存的是对应的类别编号。

3. 展平数据(可选但常用)

混淆矩阵一般是统计所有预测样本的结果,所以把二维的序列数据压成一维数组会更方便计算:

# 把序列维度展平,变成所有时间步的标签集合
true_labels_flat = true_labels.flatten()
pred_labels_flat = pred_labels.flatten()

如果你不想统计所有时间步,只想看每个序列最后一步的预测结果,那直接取每个序列的最后一个元素就行,不用展平。

4. 计算并可视化混淆矩阵

用sklearn的工具就能轻松搞定,代码示例如下:

from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay
import matplotlib.pyplot as plt

# 计算混淆矩阵
cm = confusion_matrix(true_labels_flat, pred_labels_flat)

# 可视化(按需使用)
# 记得把display_labels换成你的类别名称列表,比如["类别1", "类别2"...]
disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=你的类别名称列表)
disp.plot(cmap=plt.cm.Blues)
plt.show()

重要的小细节

如果你的模型是处理文本这类带填充(padding)的序列,记得过滤掉填充标签!比如你的填充标签是0,那就在展平后加个过滤:

# 过滤掉填充的无效标签(假设填充标签为0)
mask = true_labels_flat != 0
filtered_true = true_labels_flat[mask]
filtered_pred = pred_labels_flat[mask]
# 用过滤后的标签计算混淆矩阵
cm = confusion_matrix(filtered_true, filtered_pred)

这样统计出来的结果才是有效的,不会被填充的无效标签干扰。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 03:41:40