LIENWEIYING/CT_CTC_T1_OAR/chiasm.py

166 lines
6 KiB
Python
Raw Normal View History

2026-07-22 06:50:09 +00:00
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)}")