Download detection/filter_data.py from diing/AURAD: direct link, hf CLI and curl.
- Browser
- Download file 15.4 kB
-
https://huggingface.co/diing/AURAD/resolve/main/detection/filter_data.py
- Command line
-
hf download hf://diing/AURAD/detection/filter_data.py
-
curl -L -o filter_data.py https://huggingface.co/diing/AURAD/resolve/main/detection/filter_data.py
15.4 kB
| """ | |
| 生成图片病灶 Mask 质量筛选流程 | |
| ================================ | |
| test.json 格式(每行一个 JSON): | |
| { | |
| "file_name": "AURAD_infer/.../syn_36351.png", # 生成图片路径 | |
| "prompt": "severe Cardiomegaly on heart, ...", | |
| "attn_list": [ | |
| ["Cardiomegaly", "AURAD_dataset/.../mask_Cardiomegaly_36351.png"], | |
| ["Effusion", "AURAD_dataset/.../mask_Effusion_36351.png"], | |
| ], # GT mask 路径(PNG,白色前景) | |
| "true_label": ["Cardiomegaly", "Effusion"], | |
| "pred_label": ["Atelectasis", "Cardiomegaly", "Effusion"], | |
| "try": 1 | |
| } | |
| 筛选逻辑(4 层,逐层宽松): | |
| L1 存在性 -- Mask R-CNN 至少检测到一个目标,置信度 >= conf_thresh (default 0.25) | |
| L2 重叠性 -- 对每个疾病类别,预测 mask 与 GT mask 满足 IoU>=0.1 OR Dice>=0.15 | |
| L3 方向性 -- 质心距离 / 对角线 < centroid_ratio_thresh (default 0.45) | |
| L4 综合分 -- score = 0.4xIoU + 0.4xDice + 0.2xconf >= score_threshold (default 0.12) | |
| 输出: | |
| filtered_output.json -- 包含 high_quality / acceptable / rejected 三档 | |
| kept_list.txt -- 保留图片路径列表(按 composite_score 降序) | |
| 用法: | |
| python filter_generated_images.py \\ | |
| --test_json test.json \\ | |
| --checkpoint /home/shuhan/xray/AURA/detection/checkpoint/maskrcnn_real/checkpoint_epoch_50.pth \\ | |
| --data_root /home/shuhan/xray/AURA \\ | |
| --output_json filtered_output.json | |
| """ | |
| import os | |
| import json | |
| import argparse | |
| from pathlib import Path | |
| import numpy as np | |
| from PIL import Image | |
| from tqdm import tqdm | |
| import torch | |
| import torchvision.transforms.functional as TF | |
| from torchvision.models.detection import maskrcnn_resnet50_fpn | |
| from torchvision.models.detection.faster_rcnn import FastRCNNPredictor | |
| from torchvision.models.detection.mask_rcnn import MaskRCNNPredictor | |
| # ====================================================== | |
| # 已知疾病类别(index 0 = background,其余按训练时顺序) | |
| # checkpoint 输出层 shape=[14,...] => num_classes=14 => 13 个前景类 | |
| # 如果顺序不对请对照训练代码的 category list 修改 | |
| # ====================================================== | |
| DISEASE_CLASSES = [ | |
| "__background__", | |
| "Atelectasis", "Cardiomegaly", "Consolidation", | |
| "Effusion", "Emphysema", "Fibrosis", | |
| "Hernia", "Infiltration", "Mass", "Nodule", | |
| "Pleural_Thickening", "Pneumonia", "Pneumothorax", | |
| ] | |
| # 共 14 项(1 background + 13 前景),与 checkpoint 的 [14, 1024] 对齐 | |
| # ------------------------------------------------------ | |
| # 1. 模型构建 | |
| # ------------------------------------------------------ | |
| def build_model(checkpoint_path, device): | |
| """从 checkpoint 自动推断 num_classes,避免手动维护类别数出错""" | |
| ckpt = torch.load(checkpoint_path, map_location=device) | |
| state = ckpt.get("model", ckpt.get("state_dict", ckpt)) | |
| # cls_score.weight shape = [num_classes, 1024] | |
| num_classes = state["roi_heads.box_predictor.cls_score.weight"].shape[0] | |
| print(f" [Info] num_classes auto-detected from checkpoint: {num_classes}") | |
| model = maskrcnn_resnet50_fpn(pretrained=False) | |
| in_feat = model.roi_heads.box_predictor.cls_score.in_features | |
| model.roi_heads.box_predictor = FastRCNNPredictor(in_feat, num_classes) | |
| in_feat_mask = model.roi_heads.mask_predictor.conv5_mask.in_channels | |
| model.roi_heads.mask_predictor = MaskRCNNPredictor(in_feat_mask, 256, num_classes) | |
| missing, unexpected = model.load_state_dict(state, strict=False) | |
| if missing: | |
| print(f" [Warn] Missing keys : {missing[:3]}") | |
| if unexpected: | |
| print(f" [Warn] Unexpected keys: {unexpected[:3]}") | |
| model.to(device).eval() | |
| return model, num_classes | |
| # ------------------------------------------------------ | |
| # 2. 工具函数 | |
| # ------------------------------------------------------ | |
| def load_gt_mask(mask_path, data_root): | |
| """加载 GT mask PNG,返回二值 numpy array (H, W),文件不存在返回 None""" | |
| p = Path(mask_path) | |
| if not p.is_absolute(): | |
| p = Path(data_root) / p | |
| if not p.exists(): | |
| return None | |
| arr = np.array(Image.open(p).convert("L")) | |
| return (arr > 127).astype(np.uint8) | |
| def calc_iou(a, b): | |
| inter = np.logical_and(a, b).sum() | |
| union = np.logical_or(a, b).sum() | |
| return float(inter) / (float(union) + 1e-6) | |
| def calc_dice(a, b): | |
| inter = np.logical_and(a, b).sum() | |
| return 2.0 * float(inter) / (float(a.sum() + b.sum()) + 1e-6) | |
| def calc_centroid_ratio(pred, gt): | |
| """质心距离 / 图像对角线,越小越好,1.0 表示无法计算""" | |
| h, w = pred.shape | |
| def centroid(m): | |
| ys, xs = np.where(m > 0) | |
| return (float(xs.mean()), float(ys.mean())) if len(ys) > 0 else None | |
| cp, cg = centroid(pred), centroid(gt) | |
| if cp is None or cg is None: | |
| return 1.0 | |
| dist = ((cp[0] - cg[0]) ** 2 + (cp[1] - cg[1]) ** 2) ** 0.5 | |
| return dist / ((h ** 2 + w ** 2) ** 0.5 + 1e-6) | |
| # ------------------------------------------------------ | |
| # 3. 单张图评估 | |
| # ------------------------------------------------------ | |
| def evaluate_one_image( | |
| model, img_tensor, attn_list, data_root, | |
| disease_classes, device, | |
| conf_thresh=0.25, mask_thresh=0.5, | |
| iou_min=0.10, dice_min=0.15, centroid_max=0.45, | |
| ): | |
| """ | |
| 推理 + 与各 GT mask 对比。 | |
| 返回 {"pass": True/False, "reason": ..., "composite_score": ..., "per_class": [...]} | |
| """ | |
| outputs = model([img_tensor.to(device)])[0] | |
| pred_scores = outputs["scores"].cpu().numpy() | |
| pred_labels = outputs["labels"].cpu().numpy() | |
| pred_masks = outputs["masks"].cpu().numpy() # (N, 1, H, W) | |
| # L1: 置信度 | |
| keep = pred_scores >= conf_thresh | |
| if keep.sum() == 0: | |
| return {"pass": False, "reason": "L1_no_detection_above_conf_thresh"} | |
| pred_scores = pred_scores[keep] | |
| pred_labels = pred_labels[keep] | |
| pred_masks = pred_masks[keep] | |
| per_class = [] | |
| for disease_name, mask_path in attn_list: | |
| gt_mask = load_gt_mask(mask_path, data_root) | |
| if gt_mask is None: | |
| continue # GT mask 文件缺失,跳过该类别 | |
| # 找对应类别的预测索引(找不到则用全部预测) | |
| cls_idx = next( | |
| (i for i, c in enumerate(disease_classes) | |
| if c.lower() == disease_name.lower()), | |
| None | |
| ) | |
| if cls_idx is not None: | |
| cand_idx = np.where(pred_labels == cls_idx)[0] | |
| else: | |
| cand_idx = np.arange(len(pred_labels)) | |
| if len(cand_idx) == 0: | |
| per_class.append({ | |
| "disease": disease_name, | |
| "iou": 0.0, "dice": 0.0, "conf": 0.0, | |
| "centroid_dist_ratio": 1.0, | |
| "composite_score": 0.0, | |
| "pass_L2": False, "pass_L3": False, | |
| }) | |
| continue | |
| # 在候选预测中选综合分最高的 | |
| best = None | |
| best_cs = -1.0 | |
| for i in cand_idx: | |
| pred_bin = (pred_masks[i, 0] >= mask_thresh).astype(np.uint8) | |
| # 防止推理时自动 resize 导致尺寸不匹配 | |
| if pred_bin.shape != gt_mask.shape: | |
| pred_bin = (np.array( | |
| Image.fromarray(pred_bin * 255).resize( | |
| (gt_mask.shape[1], gt_mask.shape[0]), Image.NEAREST | |
| ) | |
| ) > 127).astype(np.uint8) | |
| iou_v = calc_iou(pred_bin, gt_mask) | |
| dice_v = calc_dice(pred_bin, gt_mask) | |
| conf_v = float(pred_scores[i]) | |
| cdr_v = calc_centroid_ratio(pred_bin, gt_mask) | |
| pass_L2 = (iou_v >= iou_min) or (dice_v >= dice_min) | |
| pass_L3 = cdr_v <= centroid_max | |
| cs = 0.4 * iou_v + 0.4 * dice_v + 0.2 * conf_v | |
| if cs > best_cs: | |
| best_cs = cs | |
| best = { | |
| "disease": disease_name, | |
| "iou": round(iou_v, 4), | |
| "dice": round(dice_v, 4), | |
| "conf": round(conf_v, 4), | |
| "centroid_dist_ratio": round(cdr_v, 4), | |
| "composite_score": round(cs, 4), | |
| "pass_L2": pass_L2, | |
| "pass_L3": pass_L3, | |
| } | |
| per_class.append(best) | |
| if not per_class: | |
| return {"pass": False, "reason": "L1_no_gt_mask_loaded"} | |
| # 汇总:只要有一个疾病类别同时通过 L2 + L3,整张图保留 | |
| # (多病灶图宽松策略:其中一个病灶定位准确即可) | |
| if not any(r["pass_L2"] for r in per_class): | |
| return {"pass": False, "reason": "L2_no_overlap_with_any_gt_mask", | |
| "per_class": per_class} | |
| if not any(r["pass_L3"] for r in per_class): | |
| return {"pass": False, "reason": "L3_all_centroids_too_far", | |
| "per_class": per_class} | |
| # 综合分取各类别最高值 | |
| final_score = max(r["composite_score"] for r in per_class) | |
| return {"pass": True, "composite_score": round(final_score, 4), | |
| "per_class": per_class} | |
| # ------------------------------------------------------ | |
| # 4. 主流程 | |
| # ------------------------------------------------------ | |
| def run_filter(args): | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| print(f"[Info] Device : {device}") | |
| print(f"[Info] Checkpoint : {args.checkpoint}") | |
| print(f"[Info] Data root : {args.data_root}") | |
| print(f"[Info] test.json : {args.test_json}") | |
| model, num_classes = build_model(args.checkpoint, device) | |
| print(f"[Info] Num classes : {num_classes} (auto from checkpoint)") | |
| print("[Info] Model ready.\n") | |
| # 读取 test.json(每行一个 JSON object / 也兼容整体 JSON list) | |
| records = [] | |
| with open(args.test_json, "r") as f: | |
| raw = f.read().strip() | |
| try: | |
| # 尝试整体解析(list 格式) | |
| parsed = json.loads(raw) | |
| records = parsed if isinstance(parsed, list) else [parsed] | |
| except json.JSONDecodeError: | |
| # 逐行解析(jsonl 格式) | |
| for line in raw.splitlines(): | |
| line = line.strip() | |
| if line: | |
| records.append(json.loads(line)) | |
| print(f"[Info] Total records : {len(records)}\n") | |
| high_quality = [] # 保留原始 rec | |
| acceptable = [] | |
| rejected = [] | |
| score_map = {} # file_name -> composite_score,用于排序 | |
| for rec in tqdm(records, desc="Filtering"): | |
| gen_path = rec["file_name"] | |
| if not os.path.isabs(gen_path): | |
| gen_path = os.path.join(args.data_root, gen_path) | |
| if not os.path.exists(gen_path): | |
| rejected.append({**rec, "reason": "generated_image_not_found"}) | |
| continue | |
| try: | |
| pil_img = Image.open(gen_path).convert("RGB") | |
| except Exception as e: | |
| rejected.append({**rec, "reason": f"image_load_error:{e}"}) | |
| continue | |
| img_tensor = TF.to_tensor(pil_img) | |
| result = evaluate_one_image( | |
| model, img_tensor, | |
| attn_list = rec["attn_list"], | |
| data_root = args.data_root, | |
| disease_classes = DISEASE_CLASSES, | |
| device = device, | |
| conf_thresh = args.conf_thresh, | |
| mask_thresh = 0.5, | |
| iou_min = args.iou_min, | |
| dice_min = args.dice_min, | |
| centroid_max = args.centroid_max, | |
| ) | |
| if not result["pass"]: | |
| rejected.append(rec) | |
| continue | |
| cs = result["composite_score"] | |
| score_map[rec["file_name"]] = cs | |
| if cs >= args.high_quality_threshold: | |
| high_quality.append(rec) | |
| elif cs >= args.score_threshold: | |
| acceptable.append(rec) | |
| else: | |
| rejected.append(rec) | |
| # 按综合分降序排列保留的图片 | |
| high_quality.sort(key=lambda x: score_map.get(x["file_name"], 0), reverse=True) | |
| acceptable.sort(key=lambda x: score_map.get(x["file_name"], 0), reverse=True) | |
| total = len(records) | |
| n_hq = len(high_quality) | |
| n_acc = len(acceptable) | |
| n_rej = len(rejected) | |
| print("\n" + "=" * 52) | |
| print(f" Total : {total}") | |
| print(f" High quality : {n_hq:5d} ({100*n_hq/total:.1f}%)") | |
| print(f" Acceptable : {n_acc:5d} ({100*n_acc/total:.1f}%)") | |
| print(f" Rejected : {n_rej:5d} ({100*n_rej/total:.1f}%)") | |
| print(f" Retention rate : {100*(n_hq+n_acc)/total:.1f}%") | |
| print("=" * 52) | |
| # 拒绝原因统计 | |
| reasons = {} | |
| for r in rejected: | |
| k = r.get("reason", "unknown") | |
| if k.startswith("L4"): | |
| k = "L4_low_composite_score" | |
| reasons[k] = reasons.get(k, 0) + 1 | |
| print("\n Rejection breakdown:") | |
| for k, v in sorted(reasons.items(), key=lambda x: -x[1]): | |
| print(f" {k:<42s}: {v}") | |
| # 输出:与输入 test.json 完全相同的格式(每行一个 JSON),只保留通过筛选的图片 | |
| kept = high_quality + acceptable | |
| with open(args.output_json, "w") as f: | |
| for rec in kept: | |
| f.write(json.dumps(rec, ensure_ascii=False) + "\n") | |
| print(f"\n[Info] Filtered jsonl -> {args.output_json} ({len(kept)} images)") | |
| # ------------------------------------------------------ | |
| # 5. CLI | |
| # ------------------------------------------------------ | |
| if __name__ == "__main__": | |
| parser = argparse.ArgumentParser( | |
| description="Filter generated X-ray images by mask localization quality", | |
| formatter_class=argparse.ArgumentDefaultsHelpFormatter, | |
| ) | |
| parser.add_argument("--test_json", required=True) | |
| parser.add_argument("--checkpoint", | |
| default="/home/shuhan/xray/AURA/detection/checkpoint/" | |
| "maskrcnn_real/checkpoint_epoch_50.pth") | |
| parser.add_argument("--data_root", default="/home/shuhan/xray/AURA", | |
| help="Root dir prepended to relative paths in test.json") | |
| parser.add_argument("--output_json", default="filtered_output.json") | |
| g = parser.add_argument_group("Thresholds") | |
| g.add_argument("--conf_thresh", type=float, default=0.2, | |
| help="L1: Mask R-CNN confidence threshold") | |
| g.add_argument("--iou_min", type=float, default=0.25, | |
| help="L2: IoU lower bound (OR with dice_min)") | |
| g.add_argument("--dice_min", type=float, default=0.25, | |
| help="L2: Dice lower bound (OR with iou_min)") | |
| g.add_argument("--centroid_max", type=float, default=0.45, | |
| help="L3: centroid-distance/diagonal upper bound") | |
| g.add_argument("--score_threshold", type=float, default=0.20, | |
| help="L4 lower bound -> acceptable tier") | |
| g.add_argument("--high_quality_threshold", type=float, default=0.55, | |
| help="L4 upper bound -> high_quality tier") | |
| args = parser.parse_args() | |
| run_filter(args) | |
| ''' | |
| python filter_data.py \ | |
| --test_json /data16T/chestx-ray/AURAD_infer/det-train-real-mask-image-total-filter/success_train_prompt_layout2image_multi_det.json \ | |
| --data_root /data16T/chestx-ray \ | |
| --output_json filtered_CXR_SD14_unet+mask_multi_disease_channel_high.json | |
| ''' |