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

如何从打印结果统计MNIST字母分类的TP/TN/FP/FN数量?

解决MNIST字母one-hot编码的分类指标统计问题

你现在的代码只是把每个样本的TP/TN/FP/FN结果打印出来,要从这些零散的输出里按标签统计数量确实太麻烦了——与其事后去整理打印内容,不如直接在代码里完成计数统计,这样效率高得多!

改进思路

我们可以先初始化一个字典,用来存储每个标签的TP、TN、FP、FN计数,然后遍历每个标签的真实值和预测值,逐个样本判断并更新计数,最后基于这些计数计算准确率、精确率、召回率和F1分数。

完整代码实现

import pandas as pd

# 假设y_true和y_pred都是维度为[88800, 26]的DataFrame,列名为字母标签(比如'A'、'B'...或'1'、'2'...)

# 1. 初始化每个标签的统计字典,存储TP/TN/FP/FN的数量
metrics = {}
for label in y_true.columns:
    metrics[label] = {'TP': 0, 'TN': 0, 'FP': 0, 'FN': 0}

# 2. 遍历每个标签,统计对应指标
for label in y_true.columns:
    # 获取当前标签的真实值和预测值数组
    true_vals = y_true[label].values
    pred_vals = y_pred[label].values
    
    # 逐个样本判断并更新计数
    for true, pred in zip(true_vals, pred_vals):
        if true == 1 and pred == 1:
            metrics[label]['TP'] += 1
        elif true == 0 and pred == 0:
            metrics[label]['TN'] += 1
        elif true == 0 and pred == 1:
            metrics[label]['FP'] += 1
        elif true == 1 and pred == 0:
            metrics[label]['FN'] += 1

# 3. 计算每个标签的准确率、精确率、召回率、F1分数
label_metrics = {}
for label in metrics:
    tp = metrics[label]['TP']
    tn = metrics[label]['TN']
    fp = metrics[label]['FP']
    fn = metrics[label]['FN']
    
    # 总样本数
    total_samples = tp + tn + fp + fn
    # 准确率:(正确预测的正样本+正确预测的负样本)/总样本数
    accuracy = (tp + tn) / total_samples if total_samples != 0 else 0.0
    # 精确率:正确预测的正样本 / 所有预测为正的样本数(避免除以0)
    precision = tp / (tp + fp) if (tp + fp) != 0 else 0.0
    # 召回率:正确预测的正样本 / 所有真实为正的样本数
    recall = tp / (tp + fn) if (tp + fn) != 0 else 0.0
    # F1分数:精确率和召回率的调和平均数
    f1_score = 2 * (precision * recall) / (precision + recall) if (precision + recall) != 0 else 0.0
    
    # 保存当前标签的所有指标
    label_metrics[label] = {
        'TP': tp,
        'TN': tn,
        'FP': fp,
        'FN': fn,
        '准确率': round(accuracy, 4),
        '精确率': round(precision, 4),
        '召回率': round(recall, 4),
        'F1分数': round(f1_score, 4)
    }

# 4. 转换成DataFrame,方便查看和导出
result_df = pd.DataFrame(label_metrics).T
print("各标签分类指标统计:")
print(result_df)

代码说明

  • 我们直接遍历y_true的列名(也就是每个字母标签),避免了你原来代码里i-1的索引错误风险;
  • 用字典存储每个标签的计数,逻辑清晰,统计效率高;
  • 加入了分母为0的判断,避免计算时出现除以0的错误;
  • 最后转换成DataFrame,能直观地看到每个标签的所有指标。

如果你一定要从之前的打印结果统计,可以把打印内容保存到一个文本文件,然后用Python读取文件并按标签分组计数,但这种方法远不如直接在代码里统计高效,不推荐哦。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 03:53:26