AURAD_dataset / chexlocalize /write_prompt.py
diing's picture
Add files using upload-large-folder tool
e4bd0ba verified
Raw
History Blame Contribute Delete
9.46 kB
import os
import json
import numpy as np
from skimage import io
from tqdm import tqdm
def define_organ_parts(mask):
"""
organ label:
- Left Lung: 60
- Right Lung: 120
- Heart: 180
- Mediastinum: 240
"""
left_lung_coords = np.where(mask == 60)
right_lung_coords = np.where(mask == 120)
if left_lung_coords[0].size > 0:
left_min, left_max = left_lung_coords[0].min(), left_lung_coords[0].max()
left_upper_boundary = left_min + (left_max - left_min) // 3
left_middle_boundary = left_min + 2 * (left_max - left_min) // 3
else:
left_upper_boundary, left_middle_boundary = 0, 0
if right_lung_coords[0].size > 0:
right_min, right_max = right_lung_coords[0].min(), right_lung_coords[0].max()
right_upper_boundary = right_min + (right_max - right_min) // 3
right_middle_boundary = right_min + 2 * (right_max - right_min) // 3
else:
right_upper_boundary, right_middle_boundary = 0, 0
organ_parts = {
"left upper lung": (mask == 60) & (np.arange(mask.shape[0])[:, None] <= left_upper_boundary),
"left middle lung": (mask == 60) & (
(np.arange(mask.shape[0])[:, None] > left_upper_boundary) &
(np.arange(mask.shape[0])[:, None] <= left_middle_boundary)
),
"left lower lung": (mask == 60) & (np.arange(mask.shape[0])[:, None] > left_middle_boundary),
"right upper lung": (mask == 120) & (np.arange(mask.shape[0])[:, None] <= right_upper_boundary),
"right middle lung": (mask == 120) & (
(np.arange(mask.shape[0])[:, None] > right_upper_boundary) &
(np.arange(mask.shape[0])[:, None] <= right_middle_boundary)
),
"right lower lung": (mask == 120) & (np.arange(mask.shape[0])[:, None] > right_middle_boundary),
"heart": (mask == 180),
"mediastinum": (mask == 240),
}
return organ_parts
def load_image(path):
image = io.imread(path)
if len(image.shape) == 3:
image = image[..., 0]
return image
def calculate_width(region_mask):
non_zero_columns = np.where(region_mask > 0)[1]
if len(non_zero_columns) == 0:
return 0
return non_zero_columns.max() - non_zero_columns.min() + 1
def process_organ_and_mask(disease, organ_path, mask_path):
"""
Return: (disease, location_label, severity)
"""
organ = load_image(organ_path)
mask = load_image(mask_path)
organ_parts = define_organ_parts(organ)
overlap_results = {}
for part, mask_part in organ_parts.items():
overlap_area = np.sum((mask_part > 0) & (mask > 0))
if overlap_area > 0:
overlap_results[part] = overlap_area
if not overlap_results:
return None
main_part = max(overlap_results, key=overlap_results.get)
if disease == "Cardiomegaly":
if main_part == "heart":
location_label = main_part
organ_width = calculate_width(organ)
mask_width = calculate_width(mask)
if organ_width == 0:
return None
ratio = mask_width / organ_width
if ratio <= 0.55:
severity = "mild"
elif ratio < 0.6:
severity = "moderate"
else:
severity = "severe"
else:
return None
if disease == "Enlarged Cardiomediastinum":
organ_width = calculate_width(organ)
mask_width = calculate_width(mask)
if organ_width == 0 or mask_width == 0:
return None
ratio = mask_width / organ_width
location_label = "heart and mediastinum"
if ratio <= 0.55:
severity = "mild"
elif 0.55 < ratio < 0.6:
severity = "moderate"
else:
severity = "severe"
return disease, location_label, severity
else:
if overlap_results[main_part] > np.sum(mask > 0) * 0.7:
location_label = main_part
severity = "mild"
else:
left_regions = {"left upper lung", "left middle lung", "left lower lung"}
right_regions = {"right upper lung", "right middle lung", "right lower lung"}
active_regions = set(overlap_results.keys())
left_lung = (organ == 60)
left_lung_area = np.sum(left_lung)
left_overlap_area = np.sum((mask > 0) & left_lung)
left_lung_ratio = left_overlap_area / left_lung_area if left_lung_area > 0 else 0.0
right_lung = (organ == 120)
right_lung_area = np.sum(right_lung)
right_overlap_area = np.sum((mask > 0) & right_lung)
right_lung_ratio = right_overlap_area / right_lung_area if right_lung_area > 0 else 0.0
left_overlap = active_regions & left_regions
right_overlap = active_regions & right_regions
if left_overlap and right_overlap:
location_label = "bilateral lung"
elif left_overlap:
location_label = "left lung"
elif right_overlap:
location_label = "right lung"
else:
location_label = None
if disease == "Support Devices":
severity = None
else:
if left_lung_ratio < 0.3 and right_lung_ratio < 0.3:
severity = "mild"
elif (
left_lung_ratio > 0.6 or
right_lung_ratio > 0.6 or
(left_lung_ratio + right_lung_ratio) > 0.6
):
severity = "severe"
else:
severity = "moderate"
return disease, location_label, severity
def normalize_disease_name(disease):
if disease == "Airspace Opacity":
return "Opacity"
if disease == "Pleural Effusion":
return "Effusion"
return disease
def process_jsonl_from_attn_list(dataset_root, input_jsonl, output_jsonl, skip_empty_prompt=False):
"""
input jsonl example:
{
"file_name": "...jpg",
"organ": "...png",
"rib": "...jpg",
"prompt": "...",
"attn_list": [["Airspace Opacity", "...jpg"]]
}
output:
- rewrite prompt from organ + disease mask
- ignore Lung Lesion
- Airspace Opacity -> Opacity
"""
num_total = 0
num_written = 0
num_skipped_empty = 0
with open(input_jsonl, "r") as fin, open(output_jsonl, "w") as fout:
for line in tqdm(fin, desc="Processing jsonl"):
line = line.strip()
if not line:
continue
num_total += 1
item = json.loads(line)
organ_rel = item["organ"]
organ_path = os.path.join(dataset_root, organ_rel)
if not os.path.exists(organ_path):
print(f"[Warning] Missing organ mask: {organ_path}")
continue
new_attn_list = []
disease_texts = []
seen_mask = set()
for disease_name, mask_rel in item.get("attn_list", []):
# ignore Lung Lesion
if disease_name == "Lung Lesion":
continue
disease_name = normalize_disease_name(disease_name)
mask_path = os.path.join(dataset_root, mask_rel)
if not os.path.exists(mask_path):
print(f"[Warning] Missing disease mask: {mask_path}")
continue
result = process_organ_and_mask(disease_name, organ_path, mask_path)
if result is None:
continue
disease_calc, loc, severity = result
if loc is not None and severity is not None:
disease_texts.append(f"{severity} {disease_calc} on {loc}")
elif loc is not None:
disease_texts.append(f"{disease_calc} on {loc}")
else:
disease_texts.append(disease_calc)
if mask_rel not in seen_mask:
new_attn_list.append([disease_name, mask_rel])
seen_mask.add(mask_rel)
new_prompt = ", ".join(disease_texts)
if skip_empty_prompt and len(new_prompt) == 0:
num_skipped_empty += 1
continue
new_item = {
"file_name": item["file_name"],
"organ": item["organ"],
"rib": item["rib"],
"prompt": new_prompt,
"attn_list": new_attn_list,
}
fout.write(json.dumps(new_item) + "\n")
num_written += 1
print(f"Total input: {num_total}")
print(f"Written: {num_written}")
if skip_empty_prompt:
print(f"Skipped empty prompt: {num_skipped_empty}")
print(f"Saved to: {output_jsonl}")
if __name__ == "__main__":
dataset_root = "/data16T/chestx-ray/AURAD_dataset"
input_jsonl = "/data16T/chestx-ray/AURAD_dataset/chexlocalize/json-layout2image/test_prompt_text2layout_multi.json"
output_jsonl = "/data16T/chestx-ray/AURAD_dataset/chexlocalize/test_prompt_layout2image_multi.json"
process_jsonl_from_attn_list(
dataset_root=dataset_root,
input_jsonl=input_jsonl,
output_jsonl=output_jsonl,
skip_empty_prompt=True, # 如果想把空prompt样本直接丢掉,就改成 True
)