LIENWEIYING/T1C/Dice/validation/compute_validation_dice.py
2026-07-22 14:50:09 +08:00

288 lines
No EOL
11 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 os
import time
import numpy as np
import SimpleITK as sitk
from tqdm import tqdm
import json
def load_nifti(file_path):
"""加載 NIfTI 文件"""
return sitk.GetArrayFromImage(sitk.ReadImage(file_path))
def compute_dice(pred, gt):
"""計算Dice係數"""
intersection = np.sum(pred & gt)
sum_pred = np.sum(pred)
sum_gt = np.sum(gt)
if sum_pred == 0 and sum_gt == 0:
return 1.0
elif sum_pred == 0 or sum_gt == 0:
return 0.0
return 2.0 * intersection / (sum_pred + sum_gt)
def separate_structures(combined_mask, structure_name=""):
"""分離左右結構"""
sitk_image = sitk.GetImageFromArray(combined_mask)
labeled_components = sitk.ConnectedComponent(sitk_image)
labeled_array = sitk.GetArrayFromImage(labeled_components)
unique, counts = np.unique(labeled_array, return_counts=True)
if len(unique) < 3:
print(f"警告:{structure_name} 未找到足夠的連通區域")
return None, None
region_sizes = [(i, count) for i, count in zip(unique[1:], counts[1:])]
sorted_regions = sorted(region_sizes, key=lambda x: x[1], reverse=True)
if len(sorted_regions) < 2:
print(f"警告:{structure_name} 未找到兩個主要區域")
return None, None
region1_mask = (labeled_array == sorted_regions[0][0])
region2_mask = (labeled_array == sorted_regions[1][0])
def get_centroid(mask):
indices = np.where(mask)
return np.mean(indices[2])
centroid1 = get_centroid(region1_mask)
centroid2 = get_centroid(region2_mask)
if centroid1 > centroid2:
right = region1_mask
left = region2_mask
else:
right = region2_mask
left = region1_mask
return right.astype(np.uint8), left.astype(np.uint8)
def evaluate_case(gt_path, pred_path, labels):
"""評估單個案例"""
print(f"\n測試文件: {os.path.basename(pred_path)}")
start_time = time.time()
gt_img = load_nifti(gt_path)
pred_img = load_nifti(pred_path)
load_time = time.time() - start_time
print(f"圖像大小: {gt_img.shape}")
print(f"加載時間: {load_time:.2f}")
results = {}
total_time = 0
# 處理眼睛
start_time = time.time()
combined_gt_eye = ((gt_img == 2) | (gt_img == 3)).astype(np.uint8)
combined_pred_eye = ((pred_img == 2) | (pred_img == 3)).astype(np.uint8)
combined_pred_eye_sitk = sitk.GetImageFromArray(combined_pred_eye)
closing_filter = sitk.BinaryMorphologicalClosingImageFilter()
closing_filter.SetKernelRadius(2)
processed_pred_eye = sitk.GetArrayFromImage(closing_filter.Execute(combined_pred_eye_sitk))
eye_dice = compute_dice(processed_pred_eye, combined_gt_eye)
eye_time = time.time() - start_time
total_time += eye_time
results['combined_eye'] = eye_dice
print(f"眼睛合併計算時間: {eye_time:.2f}秒, Dice: {eye_dice:.4f}")
right_pred_eye, left_pred_eye = separate_structures(processed_pred_eye, "眼睛")
right_gt_eye, left_gt_eye = separate_structures(combined_gt_eye, "眼睛")
# 處理視神經
start_time = time.time()
combined_gt_nerve = ((gt_img == 5) | (gt_img == 6)).astype(np.uint8)
combined_pred_nerve = ((pred_img == 5) | (pred_img == 6)).astype(np.uint8)
combined_pred_nerve_sitk = sitk.GetImageFromArray(combined_pred_nerve)
processed_pred_nerve = sitk.GetArrayFromImage(closing_filter.Execute(combined_pred_nerve_sitk))
nerve_dice = compute_dice(processed_pred_nerve, combined_gt_nerve)
nerve_time = time.time() - start_time
total_time += nerve_time
results['combined_nerve'] = nerve_dice
print(f"視神經合併計算時間: {nerve_time:.2f}秒, Dice: {nerve_dice:.4f}")
right_pred_nerve, left_pred_nerve = separate_structures(processed_pred_nerve, "視神經")
right_gt_nerve, left_gt_nerve = separate_structures(combined_gt_nerve, "視神經")
for label in labels:
start_time = time.time()
if label == 1: # Brainstem
gt_binary = (gt_img == label).astype(np.uint8)
pred_binary = (pred_img == label).astype(np.uint8)
elif label == 2 and right_pred_eye is not None: # Right Eye
gt_binary = right_gt_eye
pred_binary = right_pred_eye
elif label == 3 and left_pred_eye is not None: # Left Eye
gt_binary = left_gt_eye
pred_binary = left_pred_eye
elif label == 4: # Optic Chiasm
gt_binary = (gt_img == label).astype(np.uint8)
pred_binary = (pred_img == label).astype(np.uint8)
elif label == 5 and right_pred_nerve is not None: # Right Optic Nerve
gt_binary = right_gt_nerve
pred_binary = right_pred_nerve
elif label == 6 and left_pred_nerve is not None: # Left Optic Nerve
gt_binary = left_gt_nerve
pred_binary = left_pred_nerve
else:
continue
try:
dice = compute_dice(pred_binary, gt_binary)
compute_time = time.time() - start_time
total_time += compute_time
results[label] = dice
print(f"標籤 {label} 計算時間: {compute_time:.2f}秒, Dice: {dice:.4f}")
except Exception as e:
print(f"處理標籤 {label} 時出錯:{str(e)}")
results[label] = np.nan
continue
print(f"單個案例總計算時間: {total_time:.2f}")
return results
def evaluate_validation_fold(gt_base_folder, validation_folder, labels):
"""評估單個fold的驗證集"""
fold_results = {label: [] for label in labels}
fold_results['combined_eye'] = []
fold_results['combined_nerve'] = []
pred_files = sorted([f for f in os.listdir(validation_folder) if f.endswith('.nii.gz')])
for pred_file in tqdm(pred_files, desc=f"Processing {os.path.basename(validation_folder)}"):
patient_id = pred_file.split('.nii.gz')[0]
gt_path = os.path.join(gt_base_folder, f"{patient_id}.nii.gz")
pred_path = os.path.join(validation_folder, pred_file)
if not os.path.exists(gt_path):
print(f"找不到對應的真實標籤文件:{gt_path}")
continue
case_results = evaluate_case(gt_path, pred_path, labels)
for key in case_results:
if not np.isnan(case_results[key]):
if key in fold_results:
fold_results[key].append(case_results[key])
return fold_results
def evaluate_model(gt_folder, model_path, model_name, labels, organ_names, output_dir):
"""評估模型所有fold的驗證集"""
print(f"\n開始評估模型:{model_name}")
all_results = {label: [] for label in labels}
all_results['combined_eye'] = []
all_results['combined_nerve'] = []
fold_results = {}
for fold in range(5):
validation_folder = os.path.join(model_path, f"fold_{fold}", "validation")
if not os.path.exists(validation_folder):
print(f"找不到fold_{fold}的驗證資料夾:{validation_folder}")
continue
print(f"\n處理 {model_name} - Fold {fold} 驗證集")
results = evaluate_validation_fold(gt_folder, validation_folder, labels)
fold_results[f"fold_{fold}"] = results
for key in results:
all_results[key].extend(results[key])
summary = {}
for key in all_results:
if all_results[key]:
mean_dice = np.mean(all_results[key])
std_dice = np.std(all_results[key])
summary[key] = {
'mean_dice': float(mean_dice),
'std_dice': float(std_dice),
'n_samples': len(all_results[key])
}
else:
summary[key] = {
'mean_dice': np.nan,
'std_dice': np.nan,
'n_samples': 0
}
model_results = {
'summary': summary,
'fold_results': fold_results
}
output_path = os.path.join(output_dir, 'T1C', 'Dice', 'validation')
os.makedirs(output_path, exist_ok=True)
filename = os.path.join(output_path, f'dice_results_{model_name}.json')
with open(filename, 'w') as f:
json.dump(model_results, f, indent=4)
print(f"\n{model_name}驗證集結果已保存到 {filename}")
print(f"\n{model_name} 模型的驗證集 Dice 結果所有fold的平均值:")
for label in labels:
if label in summary:
mean = summary[label]['mean_dice']
std = summary[label]['std_dice']
n = summary[label]['n_samples']
organ_name = organ_names.get(label, f"Label_{label}")
if not np.isnan(mean):
print(f"{organ_name}: {mean:.4f} ± {std:.4f} (n={n})")
else:
print(f"{organ_name}: 無法計算 (n={n})")
print("\n合併結構的結果:")
for key in ['combined_eye', 'combined_nerve']:
if key in summary:
mean = summary[key]['mean_dice']
std = summary[key]['std_dice']
n = summary[key]['n_samples']
name = "合併眼睛" if key == 'combined_eye' else "合併視神經"
if not np.isnan(mean):
print(f"{name}: {mean:.4f} ± {std:.4f} (n={n})")
else:
print(f"{name}: 無法計算 (n={n})")
return summary
if __name__ == "__main__":
base_dir = "/mnt/1248/onlylian"
gt_folder = os.path.join(base_dir, "nnUNet/nnUNet_raw/Dataset999/labelsTr")
output_dir = base_dir
models = {
'2D': os.path.join(base_dir, "nnUNet/nnUNet_results/Dataset999/nnUNetTrainer__nnUNetPlans__2d"),
'3D_fullres': os.path.join(base_dir, "nnUNet/nnUNet_results/Dataset999/nnUNetTrainer__nnUNetPlans__3d_fullres"),
'3D_lowres': os.path.join(base_dir, "nnUNet/nnUNet_results/Dataset999/nnUNetTrainer__nnUNetPlans__3d_lowres")
}
labels = [1, 2, 3, 4, 5, 6]
organ_names = {
1: "Brainstem",
2: "Right_Eye",
3: "Left_Eye",
4: "Optic_Chiasm",
5: "Right_Optic_Nerve",
6: "Left_Optic_Nerve"
}
print("可用的模型:")
for i, model_name in enumerate(models.keys(), 1):
print(f"{i}. {model_name}")
choice = input("\n請選擇要運行的模型 (輸入數字1-3): ")
try:
model_index = int(choice) - 1
model_names = list(models.keys())
selected_model = model_names[model_index]
evaluate_model(gt_folder, models[selected_model], selected_model, labels, organ_names, output_dir)
except (ValueError, IndexError):
print("無效的選擇請輸入1-3之間的數字。")