AURAD / text2layout /T2L_post_process.py
diing's picture
Upload folder using huggingface_hub
41c8683 verified
Raw History Blame Contribute Delete
8.55 kB
from PIL import Image
import numpy as np
from PIL import Image
from scipy.ndimage import binary_erosion, binary_dilation
from skimage.morphology import disk
def find_region(generated_image, erosion_dilation_radius=5):
red_channel = generated_image[:, :, 0] # red-channel
green_channel = generated_image[:, :, 1] # green-channel
blue_channel = generated_image[:, :, 2] # blue-channel
red_region = (red_channel > 100) & (green_channel < 80) & (blue_channel < 80)
selem = disk(erosion_dilation_radius)
mask = binary_erosion(red_region, structure=selem).astype(np.uint8)
mask = binary_dilation(mask, structure=selem).astype(np.uint8)
return mask
def define_organ_parts(mask):
if isinstance(mask, Image.Image):
mask = mask.convert("L")
mask = np.array(mask)
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_lung_x_min, left_lung_x_max = left_lung_coords[1].min(), left_lung_coords[1].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, left_max = 0, 0, 0
left_lung_x_min, left_lung_x_max = 0, 0
if right_lung_coords[0].size > 0:
right_min, right_max = right_lung_coords[0].min(), right_lung_coords[0].max()
right_lung_x_min, right_lung_x_max = right_lung_coords[1].min(), right_lung_coords[1].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, right_max = 0, 0, 0
right_lung_x_min, right_lung_x_max = 0, 0
height, width = mask.shape[0], mask.shape[1]
organ_parts = {
"left upper lung": (
(np.arange(height)[:, None] <= left_upper_boundary) &
(np.arange(width) >= left_lung_x_min) &
(np.arange(width) <= left_lung_x_max)
),
"left middle lung": (
(np.arange(height)[:, None] > left_upper_boundary) &
(np.arange(height)[:, None] <= left_middle_boundary) &
(np.arange(width) >= left_lung_x_min) &
(np.arange(width) <= left_lung_x_max)
),
"left lower lung": (
(np.arange(height)[:, None] > left_middle_boundary) &
(np.arange(width) >= left_lung_x_min) &
(np.arange(width) <= left_lung_x_max)
),
"right upper lung": (
(np.arange(height)[:, None] <= right_upper_boundary) &
(np.arange(width) >= right_lung_x_min) &
(np.arange(width) <= right_lung_x_max)
),
"right middle lung": (
(np.arange(height)[:, None] > right_upper_boundary) &
(np.arange(height)[:, None] <= right_middle_boundary) &
(np.arange(width) >= right_lung_x_min) &
(np.arange(width) <= right_lung_x_max)
),
"right lower lung": (
(np.arange(height)[:, None] > right_middle_boundary) &
(np.arange(width) >= right_lung_x_min) &
(np.arange(width) <= right_lung_x_max)
),
"heart": (mask == 180),
"mediastinum": (mask == 240)
}
return organ_parts
def calculate_width(region_mask):
non_zero_columns = np.where(region_mask> 0)[1]
if len(non_zero_columns) == 0:
return 0
max_width = non_zero_columns.max() - non_zero_columns.min() + 1
return max_width
def process_organ_and_mask(disease, organ, mask):
organ_parts = define_organ_parts(organ)
if isinstance(organ, Image.Image):
organ = organ.convert("L")
organ = np.array(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 overlap_results:
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)
cardio_ratio = mask_width / organ_width
if cardio_ratio <= 0.55:
severity = "mild"
elif 0.55 < cardio_ratio < 0.6:
severity = "moderate"
elif cardio_ratio >= 0.6:
severity = "severe"
else:
location_label = None
severity = 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_overlap = active_regions & left_regions
left_lung_area = np.sum(left_lung)
left_overlap_area = np.sum((mask > 0) & (left_lung > 0))
left_lung_ratio = left_overlap_area / left_lung_area
right_lung = (organ == 120)
right_overlap = active_regions & right_regions
right_lung_area = np.sum(right_lung)
right_overlap_area = np.sum((mask > 0) & (right_lung > 0))
right_lung_ratio = right_overlap_area / right_lung_area
if left_overlap and right_overlap:
location_label = "biliteral lung"
elif left_overlap:
location_label = "left lung"
elif right_overlap:
location_label = "right lung"
else:
location_label = None
if disease == "Chest Tube" or disease == "Pacemaker":
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"
generated_prompt = f"A Chest X-ray semantic mask with {severity} {disease} on {location_label}"
return generated_prompt
def post_process(generated_image, organ, disease, prompt):
mask = find_region(generated_image)
color_map = {
"Atelectasis": (255, 0, 0),
"Calcification": (0, 255, 0),
"Cardiomegaly": (0, 0, 255),
"Consolidation": (255, 255, 0),
"Diffuse Nodule": (255, 165, 0),
"Effusion": (0, 255, 255),
"Emphysema": (255, 0, 255),
"Fibrosis": (128, 0, 128),
"Fracture": (255, 192, 203),
"Mass": (173, 255, 47),
"Nodule": (0, 128, 255),
"Pleural Thickening": (75, 0, 130),
"Pneumothorax": (255, 105, 180)
}
generated_prompt = process_organ_and_mask(disease, organ, mask)
if generated_prompt == prompt:
organ_np = np.array(organ)
color = color_map.get(disease, [0, 0, 0])
organ_np[mask == 1] = color
return Image.fromarray(mask*255), Image.fromarray(organ_np)
else:
return None, None