File size: 2,379 Bytes
41c8683
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
"""
模型表现分析脚本
用法: python analyze.py <input.csv> [-o output_dir]

输出一个 CSV:
  - 行: model_metric (如 Real_Box_IoU)
  - 列: 各个 disease,加上一列 mean (跨疾病总平均)
  - 指标: Box_IoU, Mask_IoU, Dice, mAP@0.5 (= Detected@0.5)

输出文件名: <输入文件名>_summary.csv
"""
import argparse
import os
import pandas as pd


def main():
    parser = argparse.ArgumentParser(description="计算每个模型在各疾病上的表现")
    parser.add_argument("input_csv", help="输入 CSV 文件路径")
    parser.add_argument(
        "-o", "--output-dir", default=None,
        help="输出目录(默认与输入文件同目录)",
    )
    args = parser.parse_args()

    df = pd.read_csv(args.input_csv)
    print(f"数据总览: {len(df)} 行")
    print(f"模型: {df['model'].unique().tolist()}")
    print(f"疾病: {df['disease'].unique().tolist()}")
    print(f"指标: {df['metric'].unique().tolist()}")
    print("=" * 80)

    # 透视成宽表: 每个 (model, sample_id, disease) 一行,各 metric 是列
    wide = df.pivot_table(
        index=["model", "sample_id", "disease"],
        columns="metric",
        values="value",
    ).reset_index()

    # 重命名: Detected@0.5 -> mAP@0.5
    wide = wide.rename(columns={"Detected@0.5": "mAP@0.5"})
    metrics = ["Box_IoU", "Mask_IoU", "Dice", "mAP@0.5"]

    # 1. 每个 model × disease 上各 metric 的平均
    per_disease = wide.groupby(["model", "disease"])[metrics].mean()

    # 2. 重排成: 行=model_metric, 列=disease
    #    先 stack metric 到行,再 unstack disease 到列
    table = per_disease.stack().unstack("disease")
    # 此时索引是 (model, metric),把它拼成 "model_metric"
    table.index = [f"{m}_{metric}" for m, metric in table.index]
    table.index.name = "model_metric"

    # 3. 加一列 mean: 跨疾病的平均
    table["mean"] = table.mean(axis=1)
    table = table.round(4)

    print("\n【summary 表】")
    print(table.to_string())

    # 4. 保存
    in_basename = os.path.splitext(os.path.basename(args.input_csv))[0]
    out_dir = args.output_dir or os.path.dirname(os.path.abspath(args.input_csv))
    os.makedirs(out_dir, exist_ok=True)
    out_path = os.path.join(out_dir, f"{in_basename}_summary.csv")
    table.to_csv(out_path)
    print(f"\n已保存: {out_path}")


if __name__ == "__main__":
    main()