AURAD / classification /test_resnet.py
diing's picture
Upload folder using huggingface_hub
41c8683 verified
Raw History Blame Contribute Delete
9.3 kB
"""
ResNet50 multi-label classification - test / evaluation.
输出两份 CSV:
1. <out>_per_sample.csv 每个 sample 每个疾病一行 (prob + gt)
2. <out>_per_disease.csv 每个疾病一行 (AUROC, AP, accuracy, n_pos, ...)
最后一行是 macro 平均
支持多 checkpoint, 加 "model" 列区分.
用法:
python test_resnet50_cls.py \\
--data_dir /home/jovyan \\
--test_prompt /home/jovyan/AURAD/detection/json/test_real.json \\
--checkpoints ckpt1.pth ckpt2.pth ... \\
--output_dir ./eval_out \\
--batch_size 32
"""
import argparse
import json
import os
from pathlib import Path
import numpy as np
import pandas as pd
import torch
import torch.nn as nn
import torchvision
import torchvision.transforms as T
from PIL import Image
from sklearn.metrics import roc_auc_score, average_precision_score
from torch.utils.data import DataLoader, Dataset
from tqdm import tqdm
DISEASES = [
"Atelectasis", "Calcification", "Cardiomegaly", "Consolidation",
"Diffuse Nodule", "Effusion", "Emphysema", "Fibrosis", "Fracture",
"Mass", "Nodule", "Pleural Thickening", "Pneumothorax",
]
DISEASE2IDX = {d: i for i, d in enumerate(DISEASES)}
NUM_CLASSES = len(DISEASES)
# ---------------- Dataset ----------------
class ChestXRayClsDataset(Dataset):
def __init__(self, jsonl_path, data_dir, transform=None):
self.data_dir = Path(data_dir)
self.items = []
with open(jsonl_path) as f:
for line in f:
line = line.strip()
if line:
self.items.append(json.loads(line))
self.transform = transform
def __len__(self):
return len(self.items)
def __getitem__(self, idx):
item = self.items[idx]
img = Image.open(self.data_dir / item["file_name"]).convert("L")
label = torch.zeros(NUM_CLASSES, dtype=torch.float32)
for entry in item.get("attn_list", []):
if entry[0] in DISEASE2IDX:
label[DISEASE2IDX[entry[0]]] = 1.0
if self.transform:
img = self.transform(img)
# sample_id 用 file_name (相对路径, 唯一)
return img, label, item["file_name"]
# ---------------- Model ----------------
def get_resnet50(num_classes=NUM_CLASSES, in_channels=1):
model = torchvision.models.resnet50(weights=None)
if in_channels != 3:
model.conv1 = nn.Conv2d(
in_channels, 64, kernel_size=7, stride=2, padding=3, bias=False
)
model.fc = nn.Linear(model.fc.in_features, num_classes)
return model
def get_model_name(checkpoint_path):
"""从 checkpoint 路径取一个简短的模型名."""
parent = os.path.basename(os.path.dirname(checkpoint_path))
if parent.startswith("resnet50_"):
parent = parent[len("resnet50_"):]
if not parent:
parent = os.path.splitext(os.path.basename(checkpoint_path))[0]
return parent
# ---------------- Eval ----------------
@torch.no_grad()
def evaluate_one_ckpt(model_name, ckpt_path, dataloader, device, threshold=0.5):
print(f"\n=== {model_name} ({ckpt_path}) ===")
model = get_resnet50(NUM_CLASSES, in_channels=1).to(device)
state = torch.load(ckpt_path, map_location=device)
# 兼容两种格式:完整 checkpoint dict / 纯 state_dict
if isinstance(state, dict) and "model" in state:
state = state["model"]
model.load_state_dict(state)
model.eval()
all_probs, all_labels, all_ids = [], [], []
for images, labels, ids in tqdm(dataloader, desc=model_name):
images = images.to(device, non_blocking=True)
logits = model(images)
probs = torch.sigmoid(logits).cpu().numpy()
all_probs.append(probs)
all_labels.append(labels.numpy())
all_ids.extend(list(ids))
probs = np.concatenate(all_probs, 0) # [N, C]
labels = np.concatenate(all_labels, 0) # [N, C]
# ---- per-sample 长表 ----
rows_sample = []
preds_bin = (probs >= threshold).astype(int)
for i, sid in enumerate(all_ids):
for c, dis in enumerate(DISEASES):
rows_sample.append({
"model": model_name,
"sample_id": sid,
"disease": dis,
"prob": float(probs[i, c]),
"pred": int(preds_bin[i, c]),
"gt": int(labels[i, c]),
})
# ---- per-disease 汇总 ----
rows_disease = []
aurocs, aps = [], []
for c, dis in enumerate(DISEASES):
y, p, pb = labels[:, c], probs[:, c], preds_bin[:, c]
n_pos = int(y.sum())
n_neg = int(len(y) - n_pos)
# AUROC / AP 需要两类都有
if n_pos == 0 or n_neg == 0:
auc, ap = np.nan, np.nan
else:
auc = roc_auc_score(y, p)
ap = average_precision_score(y, p)
aurocs.append(auc)
aps.append(ap)
acc = float((pb == y).mean())
# sensitivity / specificity
tp = int(((pb == 1) & (y == 1)).sum())
fn = int(((pb == 0) & (y == 1)).sum())
tn = int(((pb == 0) & (y == 0)).sum())
fp = int(((pb == 1) & (y == 0)).sum())
sens = tp / max(tp + fn, 1) if n_pos > 0 else float("nan")
spec = tn / max(tn + fp, 1) if n_neg > 0 else float("nan")
rows_disease.append({
"model": model_name,
"disease": dis,
"n_pos": n_pos,
"n_neg": n_neg,
"AUROC": auc,
"AP": ap,
"Accuracy": acc,
"Sensitivity": sens,
"Specificity": spec,
"TP": tp, "FP": fp, "TN": tn, "FN": fn,
})
# macro 平均
rows_disease.append({
"model": model_name,
"disease": "MACRO_AVG",
"n_pos": int(labels.sum()),
"n_neg": int((1 - labels).sum()),
"AUROC": float(np.nanmean([r["AUROC"] for r in rows_disease])),
"AP": float(np.nanmean([r["AP"] for r in rows_disease])),
"Accuracy": float(np.nanmean([r["Accuracy"] for r in rows_disease])),
"Sensitivity": float(np.nanmean([r["Sensitivity"] for r in rows_disease])),
"Specificity": float(np.nanmean([r["Specificity"] for r in rows_disease])),
"TP": "", "FP": "", "TN": "", "FN": "",
})
del model
torch.cuda.empty_cache()
return rows_sample, rows_disease
# ---------------- Main ----------------
def parse_args():
ap = argparse.ArgumentParser()
ap.add_argument("--data_dir", required=True)
ap.add_argument("--test_prompt", required=True)
ap.add_argument("--checkpoints", nargs="+", required=True)
ap.add_argument("--model_names", nargs="+", default=None,
help="对应每个 ckpt 的显示名, 默认从路径推断")
ap.add_argument("--output_dir", default="./eval_out")
ap.add_argument("--output_prefix", default="resnet50_eval")
ap.add_argument("--img_resolution", type=int, default=512)
ap.add_argument("--batch_size", type=int, default=32)
ap.add_argument("--num_workers", type=int, default=4)
ap.add_argument("--threshold", type=float, default=0.5,
help="二值化阈值, 用于算 Accuracy/Sens/Spec")
return ap.parse_args()
def main():
args = parse_args()
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
out_dir = Path(args.output_dir)
out_dir.mkdir(parents=True, exist_ok=True)
# 模型名
if args.model_names is None:
names = [get_model_name(c) for c in args.checkpoints]
else:
names = args.model_names
assert len(names) == len(args.checkpoints), \
"model_names 和 checkpoints 数量必须一致"
# data (验证用的 transform, 不加随机)
val_tf = T.Compose([
T.Resize((args.img_resolution, args.img_resolution)),
T.ToTensor(),
T.Normalize(mean=[0.5], std=[0.5]),
])
test_set = ChestXRayClsDataset(args.test_prompt, args.data_dir, transform=val_tf)
print(f"Test samples: {len(test_set)}")
test_loader = DataLoader(
test_set, batch_size=args.batch_size, shuffle=False,
num_workers=args.num_workers, pin_memory=True,
)
all_sample_rows = []
all_disease_rows = []
for name, ckpt in zip(names, args.checkpoints):
try:
s_rows, d_rows = evaluate_one_ckpt(
name, ckpt, test_loader, device, threshold=args.threshold
)
all_sample_rows.extend(s_rows)
all_disease_rows.extend(d_rows)
except Exception as e:
print(f"[error] {name} failed: {e}")
import traceback; traceback.print_exc()
# 写 CSV
sample_path = out_dir / f"{args.output_prefix}_per_sample.csv"
disease_path = out_dir / f"{args.output_prefix}_per_disease.csv"
pd.DataFrame(all_sample_rows).to_csv(sample_path, index=False)
pd.DataFrame(all_disease_rows).to_csv(disease_path, index=False)
print(f"\nSaved:")
print(f" - {sample_path} ({len(all_sample_rows)} rows)")
print(f" - {disease_path} ({len(all_disease_rows)} rows)")
# 打印 disease 汇总
print("\n=== Per-disease summary ===")
print(pd.DataFrame(all_disease_rows).to_string(index=False))
if __name__ == "__main__":
main()