Download classification/test_resnet.py from diing/AURAD: direct link, hf CLI and curl.
- Browser
- Download file 9.3 kB
-
https://huggingface.co/diing/AURAD/resolve/main/classification/test_resnet.py
- Command line
-
hf download hf://diing/AURAD/classification/test_resnet.py
-
curl -L -o test_resnet.py https://huggingface.co/diing/AURAD/resolve/main/classification/test_resnet.py
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 ---------------- | |
| 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() |