AURAD / detection /filter_data.py
diing's picture
Upload folder using huggingface_hub
41c8683 verified
Raw History Blame Contribute Delete
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. 单张图评估
# ------------------------------------------------------
@torch.no_grad()
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
'''