""" ResNet50 multi-label classification - test / evaluation. 输出两份 CSV: 1. _per_sample.csv 每个 sample 每个疾病一行 (prob + gt) 2. _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()