LIENWEIYING/CT_CTC_T1_OAR/chiasm.py
2026-07-22 14:50:09 +08:00

165 lines
6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import nibabel as nib
import matplotlib.pyplot as plt
import numpy as np
import os
def load_nifti(file_path):
"""載入NIfTI檔案"""
return nib.load(file_path).get_fdata()
def create_colored_mask(data):
"""創建彩色遮罩"""
color_dict = {
1: [144/255, 238/255, 144/255], # Brainstem - 綠色
2: [255/255, 218/255, 150/255], # Right_Eye - 淺黃色
3: [205/255, 170/255, 125/255], # Left_Eye - 棕色
4: [135/255, 206/255, 235/255], # Optic_Chiasm - 藍色
5: [255/255, 99/255, 71/255], # Right_Optic_Nerve - 亮珊瑚紅
6: [255/255, 160/255, 122/255] # Left_Optic_Nerve - 淺珊瑚紅
}
# 創建一個具有透明背景的遮罩
mask = np.zeros((*data.shape, 4)) # 使用4通道RGBA
for label, color in color_dict.items():
# 為每個器官設置顏色,包括完全不透明的 alpha 通道
organ_mask = (data == label)
if np.any(organ_mask): # 只處理存在的器官
mask[organ_mask] = [*color, 1.0] # RGB + alpha
return mask
def get_brain_bounds(image):
"""獲取腦部區域的邊界"""
mask = image > np.percentile(image, 1)
rows = np.any(mask, axis=1)
cols = np.any(mask, axis=0)
rmin, rmax = np.where(rows)[0][[0, -1]]
cmin, cmax = np.where(cols)[0][[0, -1]]
center_r = (rmin + rmax) // 2
center_c = (cmin + cmax) // 2
radius = int(max(rmax - rmin, cmax - cmin) * 0.55)
margin = int(radius * 0.15)
rmin = max(center_r - radius - margin, 0)
rmax = min(center_r + radius + margin, image.shape[0])
cmin = max(center_c - radius - margin, 0)
cmax = min(center_c + radius + margin, image.shape[1])
return rmin, rmax, cmin, cmax
def check_organs_in_slice(slice_data):
"""檢查切片中存在的器官標籤"""
unique_labels = set(np.unique(slice_data))
if 0 in unique_labels:
unique_labels.remove(0)
organ_dict = {
1: "Brainstem",
2: "Right_Eye",
3: "Left_Eye",
4: "Optic_Chiasm",
5: "Right_Optic_Nerve",
6: "Left_Optic_Nerve"
}
return {organ_dict[label] for label in unique_labels if label in organ_dict}
def save_slice(image_ct, image_mri, ground_truth, prediction, slice_num, save_path):
"""保存指定切片的比較圖"""
plt.clf()
fig, axes = plt.subplots(1, 2, figsize=(25, 11))
plt.subplots_adjust(wspace=0.01)
# Process CT and MRI images
img_slice_ct = np.rot90(image_ct[:, :, slice_num])
img_slice_mri = np.rot90(image_mri[:, :, slice_num])
# Get bounds from CT image
rmin, rmax, cmin, cmax = get_brain_bounds(img_slice_ct)
# 調整CT和MRI圖像的對比度
def adjust_contrast(img):
p2, p98 = np.percentile(img, (2, 98))
return np.clip((img - p2) / (p98 - p2), 0, 1)
ct_norm = adjust_contrast(img_slice_ct)
mri_norm = adjust_contrast(img_slice_mri)
# Blend images with equal weights
blended_img = 0.5 * ct_norm + 0.5 * mri_norm
# Ground Truth
gt_slice = np.rot90(ground_truth[:, :, slice_num])
gt_mask = create_colored_mask(gt_slice)
# 使用 zorder 參數來控制圖層順序
axes[0].imshow(blended_img[rmin:rmax, cmin:cmax], cmap='gray', zorder=1)
axes[0].imshow(gt_mask[rmin:rmax, cmin:cmax], zorder=2)
axes[0].set_title('Ground Truth', fontsize=32, pad=20, weight='bold')
axes[0].axis('off')
# Prediction
pred_slice = np.rot90(prediction[:, :, slice_num])
pred_mask = create_colored_mask(pred_slice)
axes[1].imshow(blended_img[rmin:rmax, cmin:cmax], cmap='gray', zorder=1)
axes[1].imshow(pred_mask[rmin:rmax, cmin:cmax], zorder=2)
axes[1].set_title('Prediction', fontsize=32, pad=20, weight='bold')
axes[1].axis('off')
plt.savefig(save_path, bbox_inches='tight', dpi=300, pad_inches=0.05)
plt.close()
print(f"已保存切片 {slice_num} 至: {save_path}")
if __name__ == "__main__":
# 載入影像檔案
image_path_1 = '/mnt/1248/onlylian/nnUNet/nnUNet_raw/Dataset221_OAR/imagesTs/OAR_082_0000.nii.gz' #CT
image_path_2 = '/mnt/1248/onlylian/nnUNet/nnUNet_raw/Dataset221_OAR/imagesTs/OAR_082_0001.nii.gz' #MRI
gt_path = '/mnt/1248/onlylian/nnUNet/nnUNet_raw/Dataset221_OAR/labelsTs/OAR_082.nii.gz'
pred_path = '/mnt/1248/onlylian/CT_CTC_T1_OAR/output_predictions/3d_fullres/fold_4/OAR_082.nii.gz'
# 設定保存路徑
save_dir = '/mnt/1248/onlylian/CT_CTC_T1_OAR'
os.makedirs(save_dir, exist_ok=True)
# 讀取影像
print("正在載入影像...")
image_ct = load_nifti(image_path_1)
image_mri = load_nifti(image_path_2)
ground_truth = load_nifti(gt_path)
prediction = load_nifti(pred_path)
print("影像載入完成")
print("\n分析切片中的器官...")
total_slices = image_ct.shape[2]
for slice_num in range(total_slices):
gt_slice = ground_truth[:, :, slice_num]
pred_slice = prediction[:, :, slice_num]
gt_organs = check_organs_in_slice(gt_slice)
pred_organs = check_organs_in_slice(pred_slice)
if gt_organs or pred_organs:
print(f"\n切片 {slice_num}:")
print(f"Ground Truth 包含: {', '.join(sorted(gt_organs))}")
print(f"Prediction 包含: {', '.join(sorted(pred_organs))}")
print("\n開始切片選擇和保存過程...")
while True:
try:
slice_num = int(input("\n請輸入要保存的切片編號 (-1 退出): "))
if slice_num == -1:
print("程序結束")
break
if 0 <= slice_num < total_slices:
save_path = os.path.join(save_dir, f'comparison_optic_chiasm_082_{slice_num:03d}.png')
save_slice(image_ct, image_mri, ground_truth, prediction, slice_num, save_path)
else:
print(f"切片編號必須在 0 到 {total_slices-1} 之間")
except ValueError:
print("請輸入有效的數字")
except Exception as e:
print(f"發生錯誤: {str(e)}")