File size: 15,402 Bytes
41c8683 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 | """
生成图片病灶 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
''' |