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