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
'''