{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":10302914,"sourceType":"datasetVersion","datasetId":6377304},{"sourceId":11202330,"sourceType":"datasetVersion","datasetId":6994311}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! pip install --upgrade kaggle # 安装 Kaggle API","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-08T01:00:54.04338Z","iopub.execute_input":"2025-04-08T01:00:54.043644Z","iopub.status.idle":"2025-04-08T01:00:58.832575Z","shell.execute_reply.started":"2025-04-08T01:00:54.043614Z","shell.execute_reply":"2025-04-08T01:00:58.83169Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T05:00:36.05567Z","iopub.execute_input":"2025-03-29T05:00:36.055989Z","iopub.status.idle":"2025-03-29T05:00:36.181392Z","shell.execute_reply.started":"2025-03-29T05:00:36.055963Z","shell.execute_reply":"2025-03-29T05:00:36.180638Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!mkdir -p ~/.kaggle\n!cp kaggle.json ~/.kaggle/\n!chmod 600 ~/.kaggle/kaggle.json\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:42:54.663848Z","iopub.execute_input":"2025-03-30T05:42:54.664188Z","iopub.status.idle":"2025-03-30T05:42:55.228177Z","shell.execute_reply.started":"2025-03-30T05:42:54.664161Z","shell.execute_reply":"2025-03-30T05:42:55.226689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!mkdir -p ~/.kaggle\n!cp /kaggle/input/kaggle-json/kaggle.json ~/.kaggle/\n!chmod 600 ~/.kaggle/kaggle.json\n!kaggle --version\n#!kaggle competitions list  # Check if you can list the available competitions or datasets","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:42:59.876629Z","iopub.execute_input":"2025-03-30T05:42:59.876984Z","iopub.status.idle":"2025-03-30T05:43:00.952586Z","shell.execute_reply.started":"2025-03-30T05:42:59.876951Z","shell.execute_reply":"2025-03-30T05:43:00.951155Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import kagglehub\n\n# Download latest version\npath = kagglehub.dataset_download(\"gaoyang142/mm-whs\")\n\nprint(\"Path to dataset files:\", path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:41:30.226082Z","iopub.execute_input":"2025-03-30T05:41:30.226411Z","iopub.status.idle":"2025-03-30T05:41:30.296942Z","shell.execute_reply.started":"2025-03-30T05:41:30.226386Z","shell.execute_reply":"2025-03-30T05:41:30.296188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/bin/bash\n!kaggle datasets download gaoyang142/mm-whs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:42:19.628767Z","iopub.execute_input":"2025-03-30T05:42:19.629151Z","iopub.status.idle":"2025-03-30T05:42:26.262522Z","shell.execute_reply.started":"2025-03-30T05:42:19.629118Z","shell.execute_reply":"2025-03-30T05:42:26.261548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"! unzip /kaggle/working/mmwhs-dataset.zip","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T04:05:47.118762Z","iopub.execute_input":"2025-03-30T04:05:47.11916Z","iopub.status.idle":"2025-03-30T04:07:52.142441Z","shell.execute_reply.started":"2025-03-30T04:05:47.119129Z","shell.execute_reply":"2025-03-30T04:07:52.141367Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm /kaggle/working/mmwhs-dataset.zip\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T04:10:03.573749Z","iopub.execute_input":"2025-03-30T04:10:03.574113Z","iopub.status.idle":"2025-03-30T04:10:03.874385Z","shell.execute_reply.started":"2025-03-30T04:10:03.574087Z","shell.execute_reply":"2025-03-30T04:10:03.873386Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# The above is all the preliminary work. Only now will we start the formal data analysis","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nimport random\nimport glob\nimport PIL\nimport shutil\nimport numpy as np\nimport pandas as pd\nimport matplotlib\nimport matplotlib.pyplot as plt\nfrom skimage import data\nfrom skimage.util import montage\nimport skimage.transform as skTrans\nfrom skimage.transform import rotate\nfrom skimage.transform import resize\nfrom PIL import Image, ImageOps\nimport nibabel as nib\nimport keras\nimport keras.backend as K\nfrom keras.callbacks import CSVLogger\nimport tensorflow as tf\nfrom tensorflow.keras.utils import plot_model\nfrom sklearn.preprocessing import MinMaxScaler\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import classification_report\nfrom tensorflow.keras.models import *\nfrom tensorflow.keras.layers import *\nfrom tensorflow.keras.optimizers import *\nfrom tensorflow.keras.callbacks import ModelCheckpoint, ReduceLROnPlateau, EarlyStopping, TensorBoard\n# from tensorflow.keras.layers.experimental import preprocessing","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:43:39.728666Z","iopub.execute_input":"2025-03-30T05:43:39.729049Z","iopub.status.idle":"2025-03-30T05:43:39.736088Z","shell.execute_reply.started":"2025-03-30T05:43:39.729017Z","shell.execute_reply":"2025-03-30T05:43:39.735083Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TRAIN_DATASET_PATH = \"/kaggle/input/mm-whs/mr_train\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:46:42.748644Z","iopub.execute_input":"2025-03-30T05:46:42.749038Z","iopub.status.idle":"2025-03-30T05:46:42.752968Z","shell.execute_reply.started":"2025-03-30T05:46:42.749006Z","shell.execute_reply":"2025-03-30T05:46:42.751942Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nold_name = TRAIN_DATASET_PATH + \"20111228_115552XHCoronaryCTACardiacs502a005.nii\"\nnew_name = TRAIN_DATASET_PATH + \"ct_train_1001_image.nii\"\n\n# renaming the file\ntry:\n    os.rename(old_name, new_name)\n    print(\"File has been re-named successfully!\")\nexcept:\n    print(\"File is already renamed!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T05:14:18.646524Z","iopub.execute_input":"2025-03-29T05:14:18.646863Z","iopub.status.idle":"2025-03-29T05:14:18.652203Z","shell.execute_reply.started":"2025-03-29T05:14:18.646837Z","shell.execute_reply":"2025-03-29T05:14:18.651111Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# load .nii file as a numpy array\ntest_image_flair = nib.load(\"/kaggle/input/mm-whs/mr_train/mr_train_1001/mr_train_1001_image.nii\").get_fdata()\nprint(\"Shape: \", test_image_flair.shape)\nprint(\"Dtype: \", test_image_flair.dtype)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:00.482864Z","iopub.execute_input":"2025-03-30T05:48:00.48352Z","iopub.status.idle":"2025-03-30T05:48:01.710226Z","shell.execute_reply.started":"2025-03-30T05:48:00.483485Z","shell.execute_reply":"2025-03-30T05:48:01.709134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Min: \", test_image_flair.min())\nprint(\"Max: \", test_image_flair.max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:04.589645Z","iopub.execute_input":"2025-03-30T05:48:04.589945Z","iopub.status.idle":"2025-03-30T05:48:04.670399Z","shell.execute_reply.started":"2025-03-30T05:48:04.589924Z","shell.execute_reply":"2025-03-30T05:48:04.669692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scaler = MinMaxScaler()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:06.763597Z","iopub.execute_input":"2025-03-30T05:48:06.763896Z","iopub.status.idle":"2025-03-30T05:48:06.767851Z","shell.execute_reply.started":"2025-03-30T05:48:06.763873Z","shell.execute_reply":"2025-03-30T05:48:06.76684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Scale the test_image_flair array and then reshape it back to its original dimensions.\n# This ensures the data is normalized/standardized for model input without altering its spatial structure.\ntest_image_flair = scaler.fit_transform(test_image_flair.reshape(-1, test_image_flair.shape[-1])).reshape(test_image_flair.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:08.839478Z","iopub.execute_input":"2025-03-30T05:48:08.839774Z","iopub.status.idle":"2025-03-30T05:48:09.792297Z","shell.execute_reply.started":"2025-03-30T05:48:08.839749Z","shell.execute_reply":"2025-03-30T05:48:09.791498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Min: \", test_image_flair.min())\nprint(\"Max: \", test_image_flair.max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:10.838573Z","iopub.execute_input":"2025-03-30T05:48:10.838876Z","iopub.status.idle":"2025-03-30T05:48:10.930981Z","shell.execute_reply.started":"2025-03-30T05:48:10.838854Z","shell.execute_reply":"2025-03-30T05:48:10.929628Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# rescaling t1\ntest_image_t1 = nib.load('/kaggle/input/mm-whs/mr_train/mr_train_1001/mr_train_1001_label.nii').get_fdata()\n# test_image_t1 = scaler.fit_transform(test_image_t1.reshape(-1, test_image_t1.shape[-1])).reshape(test_image_t1.shape)\n# 因为是label我就当做mask一样处理了.","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:27.80462Z","iopub.execute_input":"2025-03-30T05:48:27.804912Z","iopub.status.idle":"2025-03-30T05:48:28.199041Z","shell.execute_reply.started":"2025-03-30T05:48:27.80489Z","shell.execute_reply":"2025-03-30T05:48:28.198251Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"slice = 95\n\nprint(\"Slice Number: \" + str(slice))\n\nplt.figure(figsize=(12, 8))\n\n# T1\nplt.subplot(2, 3, 1)\nplt.imshow(test_image_flair[:,:,slice])\n# plt.imshow(test_image_flair[:,:,slice], cmap='gray')\nplt.title('image')\n\n\n# T1ce\nplt.subplot(2, 3, 2)\nplt.imshow(test_image_t1[:,:,slice])\nplt.title('mask')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:31.044298Z","iopub.execute_input":"2025-03-30T05:48:31.044635Z","iopub.status.idle":"2025-03-30T05:48:31.479676Z","shell.execute_reply.started":"2025-03-30T05:48:31.044604Z","shell.execute_reply":"2025-03-30T05:48:31.478809Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Modality shape\nprint(\"image: \", test_image_flair.shape)\n\n# Segmentation shape\nprint(\"mask: \", test_image_t1.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:35.277835Z","iopub.execute_input":"2025-03-30T05:48:35.278203Z","iopub.status.idle":"2025-03-30T05:48:35.283131Z","shell.execute_reply.started":"2025-03-30T05:48:35.278171Z","shell.execute_reply":"2025-03-30T05:48:35.282242Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"slice = 95\n\nprint(\"Slice number: \" + str(slice))\n\nplt.figure(figsize=(12, 8))\n\n# Apply a 90° rotation with an automatic resizing, otherwise the display is less obvious to analyze\n# T1 - Transverse View\nplt.subplot(1, 3, 1)\nplt.imshow(test_image_flair[:,:,slice], cmap='gray')\nplt.title('T1 - Transverse View')\n\n# T1 - Frontal View\nplt.subplot(1, 3, 2)\nplt.imshow(rotate(test_image_flair[:,slice,:], 90, resize=True), cmap='gray')\nplt.title('T1 - Frontal View')\n\n# T1 - Sagittal View\nplt.subplot(1, 3, 3)\nplt.imshow(rotate(test_image_flair[slice,:,:], 90, resize=True), cmap='gray')\nplt.title('T1 - Sagittal View')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:36.938237Z","iopub.execute_input":"2025-03-30T05:48:36.938518Z","iopub.status.idle":"2025-03-30T05:48:37.435136Z","shell.execute_reply.started":"2025-03-30T05:48:36.938497Z","shell.execute_reply":"2025-03-30T05:48:37.434298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Skip 50:-50 slices since there is not much to see\nplt.figure(figsize=(10, 10))\nplt.subplot(1, 1, 1)\nplt.imshow(rotate(montage(test_image_flair[50:-50,:,:]), 90, resize=True), cmap ='gray');","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:40.756024Z","iopub.execute_input":"2025-03-30T05:48:40.756344Z","iopub.status.idle":"2025-03-30T05:48:46.17165Z","shell.execute_reply.started":"2025-03-30T05:48:40.75632Z","shell.execute_reply":"2025-03-30T05:48:46.170697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Skip 50:-50 slices since there is not much to see\nplt.figure(figsize=(10, 10))\nplt.subplot(1, 1, 1)\nplt.imshow(rotate(montage(test_image_t1[50:-50,:,:]), 90, resize=True), cmap ='gray');","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:46.172898Z","iopub.execute_input":"2025-03-30T05:48:46.173225Z","iopub.status.idle":"2025-03-30T05:48:51.858644Z","shell.execute_reply.started":"2025-03-30T05:48:46.173194Z","shell.execute_reply":"2025-03-30T05:48:51.857754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plotting the segmantation\ncmap = matplotlib.colors.ListedColormap(['#440054', '#3b528b', '#18b880', '#e6d74f'])\n#norm = matplotlib.colors.BoundaryNorm([0,200,400,600,850], cmap.N)\nnorm = matplotlib.colors.BoundaryNorm([0,0.2,1,2,10], cmap.N)\n\n# plotting the 95th slice\nplt.imshow(test_image_t1[:,:,95], cmap=cmap, norm=norm)\nplt.colorbar()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:51.860326Z","iopub.execute_input":"2025-03-30T05:48:51.860657Z","iopub.status.idle":"2025-03-30T05:48:52.083362Z","shell.execute_reply.started":"2025-03-30T05:48:51.860632Z","shell.execute_reply":"2025-03-30T05:48:52.082438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Min:\", test_image_t1.min())\nprint(\"Max:\", test_image_t1.max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:52.084484Z","iopub.execute_input":"2025-03-30T05:48:52.08483Z","iopub.status.idle":"2025-03-30T05:48:52.16382Z","shell.execute_reply.started":"2025-03-30T05:48:52.084786Z","shell.execute_reply":"2025-03-30T05:48:52.163089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Isolation of class 0\nseg_0 = test_image_t1.copy()\nseg_0[seg_0 != 0] = np.nan\n\n# Isolation of class 205 (for class 1)\nseg_1 = test_image_t1.copy()\nseg_1[seg_1 != 205] = np.nan\n\n# Isolation of class 420 (for class 2)\nseg_2 = test_image_t1.copy()\nseg_2[seg_2 != 420] = np.nan\n\n# Isolation of class 550 (for class 4)\nseg_4 = test_image_t1.copy()\nseg_4[seg_4 != 550] = np.nan\n\n# Define legend\nclass_names = ['class 0', 'Non-Enhancing Tumor (class 205)', \n               'Edema (class 420)', 'Enhancing Tumor (class 550)']\nlegend = [plt.Rectangle((0, 0), 1, 1, color=cmap(i), label=class_names[i]) for i in range(len(class_names))]\n\nfig, ax = plt.subplots(1, 5, figsize=(20, 20))\n\n\nax[0].imshow(test_image_t1[:,:,slice])\nax[0].set_title('Original Segmentation')\nax[0].legend(handles=legend, loc='lower left')\n\nax[1].imshow(seg_0[:,:,slice], cmap=cmap, norm=norm)\nax[1].set_title('Not Tumor (class 0)')\n\nax[2].imshow(seg_1[:,:,slice], cmap=cmap, norm=norm)\nax[2].set_title('Non-Enhancing Tumor (class 205)')\n\nax[3].imshow(seg_2[:,:,slice], cmap=cmap, norm=norm)\nax[3].set_title('Edema (class 420)')\n\nax[4].imshow(seg_4[:,:,slice], cmap=cmap, norm=norm)\nax[4].set_title('Enhancing Tumor (class 550)')\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:48:57.509332Z","iopub.execute_input":"2025-03-30T05:48:57.509654Z","iopub.status.idle":"2025-03-30T05:49:00.7996Z","shell.execute_reply.started":"2025-03-30T05:48:57.509623Z","shell.execute_reply":"2025-03-30T05:49:00.798605Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Unique values in test_image_t1:\", np.unique(test_image_t1))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T05:44:48.406424Z","iopub.execute_input":"2025-03-29T05:44:48.406726Z","iopub.status.idle":"2025-03-29T05:44:53.866062Z","shell.execute_reply.started":"2025-03-29T05:44:48.406704Z","shell.execute_reply":"2025-03-29T05:44:53.865252Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"seg_0 unique values: {np.unique(seg_0)}\")\nprint(f\"seg_1 unique values: {np.unique(seg_1)}\")\nprint(f\"seg_2 unique values: {np.unique(seg_2)}\")\nprint(f\"seg_4 unique values: {np.unique(seg_4)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T05:45:25.06451Z","iopub.execute_input":"2025-03-29T05:45:25.064828Z","iopub.status.idle":"2025-03-29T05:45:31.05746Z","shell.execute_reply.started":"2025-03-29T05:45:25.064804Z","shell.execute_reply":"2025-03-29T05:45:31.056568Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.colors\n\n# Get all unique values\nunique_values = np.unique(test_image_t1)\nprint(\"Unique values in test_image_t1:\", unique_values)\n\n# Create segmented images for each category\nseg_images = {}\nfor value in unique_values:\n    seg_images[value] = np.where(test_image_t1 == value, test_image_t1, np.nan)\n\n# Assuming there are enough colors to represent all categories\ncmap = matplotlib.colors.ListedColormap([\n    '#FF5733',  # Bright orange\n    '#33FF57',  # Vivid green\n    '#3357FF',  # Bright blue\n    '#FFC300',  # Bright yellow\n    '#DAF7A6',  # Light green\n    '#C70039',  # Bright red\n    '#900C3F',  # Deep burgundy\n    '#581845',  # Dark purple\n    '#FFC0CB',  # Pink\n    '#FF4500'   # Orange-red\n])\n# Update legend\nclass_names = [f'class {int(value)}' for value in unique_values]\nlegend = [plt.Rectangle((0, 0), 1, 1, color=cmap(i), label=class_names[i]) for i in range(len(class_names))]\n\n# Create subplots in two rows\nn_classes = len(unique_values)\nfig, ax = plt.subplots(2, n_classes // 2 + n_classes % 2, figsize=(25, 25))  # Calculate if an extra column is needed\n\n# Plot all segmented images\nfor i, value in enumerate(unique_values):\n    row = i // (n_classes // 2 + n_classes % 2)  # Determine the row\n    col = i % (n_classes // 2 + n_classes % 2)   # Determine the column\n    ax[row, col].imshow(seg_images[value][:, :, slice], cmap=cmap, vmin=0, vmax=np.max(unique_values))\n    ax[row, col].set_title(f'Class {int(value)}')\n    ax[row, col].legend(handles=legend, loc='lower left')\n\n# Remove any extra subplots\nfor j in range(i + 1, ax.size):\n    fig.delaxes(ax.flatten()[j])\n\nplt.tight_layout()  # Automatically adjust subplot parameters to fill the entire image area\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-29T05:59:03.024242Z","iopub.execute_input":"2025-03-29T05:59:03.024541Z","iopub.status.idle":"2025-03-29T05:59:12.976207Z","shell.execute_reply.started":"2025-03-29T05:59:03.024518Z","shell.execute_reply":"2025-03-29T05:59:12.975304Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Next comes the division of the data set","metadata":{}},{"cell_type":"code","source":"# lists of directories with studies\ntrain_and_val_directories = [f.path for f in os.scandir(TRAIN_DATASET_PATH) if f.is_dir()]\n\ndef pathListIntoIds(dirList):\n    x = []\n    for i in range(0,len(dirList)):\n        x.append(dirList[i][dirList[i].rfind('/')+1:])\n    return x\n\ntrain_and_test_ids = pathListIntoIds(train_and_val_directories);\n\ntrain_test_ids, val_ids = train_test_split(train_and_test_ids,test_size=0.2)\ntrain_ids, test_ids = train_test_split(train_test_ids,test_size=0.15)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:50:39.568752Z","iopub.execute_input":"2025-03-30T05:50:39.569121Z","iopub.status.idle":"2025-03-30T05:50:39.579622Z","shell.execute_reply.started":"2025-03-30T05:50:39.569091Z","shell.execute_reply":"2025-03-30T05:50:39.578751Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Print data distribution (Train: 68%, Test: 12%, Val: 20%)\nprint(f\"Train length: {len(train_ids)}\")\nprint(f\"Validation length: {len(val_ids)}\")\nprint(f\"Test length: {len(test_ids)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:50:52.166369Z","iopub.execute_input":"2025-03-30T05:50:52.16668Z","iopub.status.idle":"2025-03-30T05:50:52.171791Z","shell.execute_reply.started":"2025-03-30T05:50:52.166655Z","shell.execute_reply":"2025-03-30T05:50:52.171022Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Data Generator","metadata":{}},{"cell_type":"markdown","source":"First, check the size of the original image","metadata":{}},{"cell_type":"code","source":"import nibabel as nib\n\nfor i in train_ids:\n    data_path = os.path.join(TRAIN_DATASET_PATH, i, f'{i}_image.nii')\n    image = nib.load(data_path).get_fdata()\n    print(f\"Image shape for ID {i}: {image.shape}\")\n    \n    data_path = os.path.join(TRAIN_DATASET_PATH, i, f'{i}_label.nii')\n    label = nib.load(data_path).get_fdata()\n    print(f\"Label shape for ID {i}: {label.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:57:13.520168Z","iopub.execute_input":"2025-03-30T05:57:13.520465Z","iopub.status.idle":"2025-03-30T05:57:14.76643Z","shell.execute_reply.started":"2025-03-30T05:57:13.520444Z","shell.execute_reply":"2025-03-30T05:57:14.765618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define seg-areas\nSEGMENT_CLASSES = {\n    0 : 'NOT tumor',\n    1 : 'NECROTIC/CORE', # or NON-ENHANCING tumor CORE\n    2 : 'EDEMA',\n    3 : 'ENHANCING' # original 4 -> converted into 3\n}\n\n# Select Slices and Image Size\nVOLUME_SLICES = 100\nVOLUME_START_AT = 22 # first slice of volume that we will include\nIMG_SIZE=128","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:50:58.087331Z","iopub.execute_input":"2025-03-30T05:50:58.087653Z","iopub.status.idle":"2025-03-30T05:50:58.091966Z","shell.execute_reply.started":"2025-03-30T05:50:58.087616Z","shell.execute_reply":"2025-03-30T05:50:58.090965Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DataGenerator(keras.utils.Sequence):\n    'Generates data for Keras'\n    def __init__(self, list_IDs, dim=(IMG_SIZE,IMG_SIZE), batch_size = 1, n_channels = 2, shuffle=True):\n        'Initialization'\n        self.dim = dim\n        self.batch_size = batch_size\n        self.list_IDs = list_IDs\n        self.n_channels = n_channels\n        self.shuffle = shuffle\n        self.on_epoch_end()\n\n    def __len__(self):\n        'Denotes the number of batches per epoch'\n        return int(np.floor(len(self.list_IDs) / self.batch_size))\n\n    def __getitem__(self, index):\n        'Generate one batch of data'\n        # Generate indexes of the batch\n        indexes = self.indexes[index*self.batch_size:(index+1)*self.batch_size]\n\n        # Find list of IDs\n        Batch_ids = [self.list_IDs[k] for k in indexes]\n\n        # Generate data\n        X, y = self.__data_generation(Batch_ids)\n\n        return X, y\n\n    def on_epoch_end(self):\n        'Updates indexes after each epoch'\n        self.indexes = np.arange(len(self.list_IDs))\n        if self.shuffle == True:\n            np.random.shuffle(self.indexes)\n\n    def __data_generation(self, Batch_ids):\n        # Initialization (With dynamic resizing)\n        X = []\n        y = []\n    \n        for c, i in enumerate(Batch_ids):\n            case_path = os.path.join(TRAIN_DATASET_PATH, i)\n    \n            data_path = os.path.join(case_path, f'{i}_image.nii')\n            image = nib.load(data_path).get_fdata()\n    \n            data_path = os.path.join(case_path, f'{i}_label.nii')\n            label = nib.load(data_path).get_fdata()\n    \n            for j in range(VOLUME_SLICES):\n                img_slice = image[:, :, j + VOLUME_START_AT]\n                label_slice = label[:, :, j + VOLUME_START_AT]\n    \n                # Resize both image and label dynamically to fit model input\n                resized_image = cv2.resize(img_slice, (IMG_SIZE, IMG_SIZE))\n                resized_label = cv2.resize(label_slice, (IMG_SIZE, IMG_SIZE), interpolation=cv2.INTER_NEAREST)\n                X.append(resized_image)\n                y.append(resized_label)\n    \n        X = np.array(X)\n        y = np.array(y)\n    \n        # Perform normalization and other operations on X and y\n        return X / np.max(X), y\n            \ntraining_generator = DataGenerator(train_ids)\nvalid_generator = DataGenerator(val_ids)\ntest_generator = DataGenerator(test_ids)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:57:39.458222Z","iopub.execute_input":"2025-03-30T05:57:39.458569Z","iopub.status.idle":"2025-03-30T05:57:39.468237Z","shell.execute_reply.started":"2025-03-30T05:57:39.458544Z","shell.execute_reply":"2025-03-30T05:57:39.467082Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define a function to display one slice and its segmentation\ndef display_slice_and_segmentation(flair, t1ce, segmentation):\n    fig, axes = plt.subplots(1, 3, figsize=(10, 5))\n\n    axes[0].imshow(flair, cmap='gray')\n    axes[0].set_title('Flair')\n    axes[0].axis('off')\n\n    axes[1].imshow(t1ce, cmap='gray')\n    axes[1].set_title('T1CE')\n    axes[1].axis('off')\n\n    axes[2].imshow(segmentation) # Displaying segmentation\n    axes[2].set_title('Segmentation')\n    axes[2].axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\n\n# Retrieve the batch from the training generator\nX_batch, Y_batch = training_generator[8]\n\n# Extract Flair, T1CE, and segmentation from the batch\nflair_batch = X_batch[:, :, :, 0]\nt1ce_batch = X_batch[:, :, :, 1]\nsegmentation_batch = np.argmax(Y_batch, axis=-1)  # Convert one-hot encoded to categorical\n\n# Extract the 50th slice from Flair, T1CE, and segmentation\nslice_index = 60  # Indexing starts from 0\nslice_flair = flair_batch[slice_index]\nslice_t1ce = t1ce_batch[slice_index]\nslice_segmentation = segmentation_batch[slice_index]\n\n# Display the 50th slice and its segmentation\ndisplay_slice_and_segmentation(slice_flair, slice_t1ce, slice_segmentation)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:57:40.682894Z","iopub.execute_input":"2025-03-30T05:57:40.683241Z","iopub.status.idle":"2025-03-30T05:57:40.843667Z","shell.execute_reply.started":"2025-03-30T05:57:40.68321Z","shell.execute_reply":"2025-03-30T05:57:40.842538Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(training_generator))\nprint(f'Input shape: {x.shape}, Output shape: {y.shape}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T04:55:48.111789Z","iopub.execute_input":"2025-03-30T04:55:48.112138Z","iopub.status.idle":"2025-03-30T04:55:49.692631Z","shell.execute_reply.started":"2025-03-30T04:55:48.112109Z","shell.execute_reply":"2025-03-30T04:55:49.691714Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Temporary test to validate the generator\nsample_X, sample_y = valid_generator[0]\nprint(\"Validation data shape:\", sample_X.shape)  # Should be (batch_size, 64, 64, 2)\nprint(\"Validation labels shape:\", sample_y.shape)  # Should be (batch_size, 64, 64, 6) (after one-hot encoding)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:02:45.90708Z","iopub.execute_input":"2025-03-30T05:02:45.907432Z","iopub.status.idle":"2025-03-30T05:02:47.792676Z","shell.execute_reply.started":"2025-03-30T05:02:45.907406Z","shell.execute_reply":"2025-03-30T05:02:47.791457Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. loss","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\n\ndef dice_coef(y_true, y_pred, smooth=1e-5):\n    \"\"\"\n    Multi-class Dice coefficient loss function (supports class imbalance)\n    \"\"\"\n    y_true = tf.cast(y_true, tf.float32)\n    y_pred = tf.math.sigmoid(y_pred)  # Apply sigmoid if the output layer is not activated\n    \n    # Calculate Dice coefficient for each class\n    intersection = tf.reduce_sum(y_true * y_pred, axis=[1, 2])\n    union = tf.reduce_sum(y_true, axis=[1, 2]) + tf.reduce_sum(y_pred, axis=[1, 2])\n    dice = (2. * intersection + smooth) / (union + smooth)\n    \n    # Ignore background (class 0)\n    dice = dice[:, 1:]  # Assuming class 0 is the background\n    \n    # Average over batches and classes\n    return tf.reduce_mean(dice)\n\ndef dice_loss(y_true, y_pred):\n    return 1 - dice_coef(y_true, y_pred)\n\n# Single class Dice coefficient (example: left ventricle, class 1)\ndef dice_coef_left_ventricle(y_true, y_pred, smooth=1e-5):\n    y_true = y_true[..., 1]  # Left ventricle corresponds to index 1\n    y_pred = y_pred[..., 1]\n    intersection = tf.reduce_sum(y_true * y_pred)\n    union = tf.reduce_sum(y_true) + tf.reduce_sum(y_pred)\n    return (2. * intersection + smooth) / (union + smooth)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:08:37.102512Z","iopub.execute_input":"2025-03-30T05:08:37.102806Z","iopub.status.idle":"2025-03-30T05:08:37.108963Z","shell.execute_reply.started":"2025-03-30T05:08:37.102783Z","shell.execute_reply":"2025-03-30T05:08:37.108187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define per class evaluation of dice coef\ndef dice_coef_necrotic(y_true, y_pred, epsilon=1e-6):\n    intersection = K.sum(K.abs(y_true[:,:,:,1] * y_pred[:,:,:,1]))\n    return (2. * intersection) / (K.sum(K.square(y_true[:,:,:,1])) + K.sum(K.square(y_pred[:,:,:,1])) + epsilon)\n\ndef dice_coef_edema(y_true, y_pred, epsilon=1e-6):\n    intersection = K.sum(K.abs(y_true[:,:,:,2] * y_pred[:,:,:,2]))\n    return (2. * intersection) / (K.sum(K.square(y_true[:,:,:,2])) + K.sum(K.square(y_pred[:,:,:,2])) + epsilon)\n\ndef dice_coef_enhancing(y_true, y_pred, epsilon=1e-6):\n    intersection = K.sum(K.abs(y_true[:,:,:,3] * y_pred[:,:,:,3]))\n    return (2. * intersection) / (K.sum(K.square(y_true[:,:,:,3])) + K.sum(K.square(y_pred[:,:,:,3])) + epsilon)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T04:10:27.760716Z","iopub.execute_input":"2025-03-30T04:10:27.760968Z","iopub.status.idle":"2025-03-30T04:10:27.768258Z","shell.execute_reply.started":"2025-03-30T04:10:27.760948Z","shell.execute_reply":"2025-03-30T04:10:27.767588Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Computing Precision\ndef precision(y_true, y_pred):\n    true_positives = tf.keras.backend.sum(tf.keras.backend.round(tf.keras.backend.clip(y_true * y_pred, 0, 1)))\n    predicted_positives = tf.keras.backend.sum(tf.keras.backend.round(tf.keras.backend.clip(y_pred, 0, 1)))\n    precision_value = true_positives / (predicted_positives + tf.keras.backend.epsilon())\n    return precision_value\n\n# Computing Sensitivity\ndef sensitivity(y_true, y_pred):\n    true_positives = K.sum(K.round(K.clip(y_true * y_pred, 0, 1)))\n    possible_positives = K.sum(K.round(K.clip(y_true, 0, 1)))\n    return true_positives / (possible_positives + K.epsilon())\n\n\n# Computing Specificity\ndef specificity(y_true, y_pred):\n    true_negatives = K.sum(K.round(K.clip((1-y_true) * (1-y_pred), 0, 1)))\n    possible_negatives = K.sum(K.round(K.clip(1-y_true, 0, 1)))\n    return true_negatives / (possible_negatives + K.epsilon())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T04:10:28.417181Z","iopub.execute_input":"2025-03-30T04:10:28.417458Z","iopub.status.idle":"2025-03-30T04:10:28.423444Z","shell.execute_reply.started":"2025-03-30T04:10:28.417435Z","shell.execute_reply":"2025-03-30T04:10:28.422552Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 7. Define the model","metadata":{}},{"cell_type":"code","source":"from keras.models import Model\nfrom keras.layers import Input, Conv2D, MaxPooling2D, UpSampling2D, concatenate, Dropout\n\ndef build_unet(input_layer, ker_init='he_normal', dropout_rate=0.2):\n    # Downsampling path\n    conv1 = Conv2D(32, 3, activation='relu', padding='same', kernel_initializer=ker_init)(input_layer)\n    conv1 = Conv2D(32, 3, activation='relu', padding='same', kernel_initializer=ker_init)(conv1)\n    pool1 = MaxPooling2D(pool_size=(2, 2))(conv1)\n\n    conv2 = Conv2D(64, 3, activation='relu', padding='same', kernel_initializer=ker_init)(pool1)\n    conv2 = Conv2D(64, 3, activation='relu', padding='same', kernel_initializer=ker_init)(conv2)\n    pool2 = MaxPooling2D(pool_size=(2, 2))(conv2)\n\n    conv3 = Conv2D(128, 3, activation='relu', padding='same', kernel_initializer=ker_init)(pool2)\n    conv3 = Conv2D(128, 3, activation='relu', padding='same', kernel_initializer=ker_init)(conv3)\n    pool3 = MaxPooling2D(pool_size=(2, 2))(conv3)\n\n    conv4 = Conv2D(256, 3, activation='relu', padding='same', kernel_initializer=ker_init)(pool3)\n    conv4 = Conv2D(256, 3, activation='relu', padding='same', kernel_initializer=ker_init)(conv4)\n    drop4 = Dropout(dropout_rate)(conv4)\n    pool4 = MaxPooling2D(pool_size=(2, 2))(drop4)\n\n    conv5 = Conv2D(512, 3, activation='relu', padding='same', kernel_initializer=ker_init)(pool4)\n    conv5 = Conv2D(512, 3, activation='relu', padding='same', kernel_initializer=ker_init)(conv5)\n    drop5 = Dropout(dropout_rate)(conv5)\n\n    # Upsampling path\n    up6 = Conv2D(256, 2, activation='relu', padding='same', kernel_initializer=ker_init)(UpSampling2D(size=(2, 2))(drop5))\n    merge6 = concatenate([drop4, up6], axis=3)\n    conv6 = Conv2D(256, 3, activation='relu', padding='same', kernel_initializer=ker_init)(merge6)\n    conv6 = Conv2D(256, 3, activation='relu', padding='same', kernel_initializer=ker_init)(conv6)\n\n    up7 = Conv2D(128, 2, activation='relu', padding='same', kernel_initializer=ker_init)(UpSampling2D(size=(2, 2))(conv6))\n    merge7 = concatenate([conv3, up7], axis=3)\n    conv7 = Conv2D(128, 3, activation='relu', padding='same', kernel_initializer=ker_init)(merge7)\n    conv7 = Conv2D(128, 3, activation='relu', padding='same', kernel_initializer=ker_init)(conv7)\n\n    up8 = Conv2D(64, 2, activation='relu', padding='same', kernel_initializer=ker_init)(UpSampling2D(size=(2, 2))(conv7))\n    merge8 = concatenate([conv2, up8], axis=3)\n    conv8 = Conv2D(64, 3, activation='relu', padding='same', kernel_initializer=ker_init)(merge8)\n    conv8 = Conv2D(64, 3, activation='relu', padding='same', kernel_initializer=ker_init)(conv8)\n\n    up9 = Conv2D(32, 2, activation='relu', padding='same', kernel_initializer=ker_init)(UpSampling2D(size=(2, 2))(conv8))\n    merge9 = concatenate([conv1, up9], axis=3)\n    conv9 = Conv2D(32, 3, activation='relu', padding='same', kernel_initializer=ker_init)(merge9)\n    conv9 = Conv2D(32, 3, activation='relu', padding='same', kernel_initializer=ker_init)(conv9)\n\n    # Final convolution layer using 1x1 convolution to map to the number of classes\n    output_layer = Conv2D(len(SEGMENT_CLASSES), (1, 1), activation='softmax')(conv9)\n\n    # Create the model\n    model = Model(inputs=input_layer, outputs=output_layer)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T04:10:31.066862Z","iopub.execute_input":"2025-03-30T04:10:31.067248Z","iopub.status.idle":"2025-03-30T04:10:31.333935Z","shell.execute_reply.started":"2025-03-30T04:10:31.067215Z","shell.execute_reply":"2025-03-30T04:10:31.333071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define input\ninput_layer = Input((IMG_SIZE, IMG_SIZE, 2))\n\n# Create model\nmodel = build_unet(input_layer, 'he_normal', 0.2)\n\n# Compile model\nmodel.compile(loss=dice_loss,\n              optimizer=keras.optimizers.Adam(learning_rate=1e-4),\n              metrics=[\n                  'accuracy',\n                  tf.keras.metrics.MeanIoU(num_classes=len(SEGMENT_CLASSES)), \n                  dice_coef, \n                  dice_coef_left_ventricle,\n                  keras.metrics.CategoricalAccuracy(name='accuracy'),\n                  keras.metrics.MeanIoU(num_classes=len(SEGMENT_CLASSES), name='mean_io_u'),\n                  precision, \n                  sensitivity, \n                  specificity, \n                  dice_coef_necrotic, \n                  dice_coef_edema, \n                  dice_coef_enhancing\n              ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:09:46.511872Z","iopub.execute_input":"2025-03-30T05:09:46.512206Z","iopub.status.idle":"2025-03-30T05:09:46.696429Z","shell.execute_reply.started":"2025-03-30T05:09:46.512178Z","shell.execute_reply":"2025-03-30T05:09:46.695618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check and verify the label distribution\nfor i in range(len(val_ids)):\n    label_path = os.path.join(TRAIN_DATASET_PATH, f'mr_train_{i}_label.nii')\n    label = nib.load(label_path).get_fdata()\n    print(f\"验证文件 {i} 的标签分布:\", np.unique(label, return_counts=True))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:14:52.875236Z","iopub.execute_input":"2025-03-30T05:14:52.875548Z","iopub.status.idle":"2025-03-30T05:14:52.894631Z","shell.execute_reply.started":"2025-03-30T05:14:52.875519Z","shell.execute_reply":"2025-03-30T05:14:52.893398Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_model(model,\n           show_shapes = True,\n           show_dtype=False,\n           show_layer_names = True,\n           rankdir = 'TB',\n           expand_nested = False,\n           dpi = 70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T04:10:35.617122Z","iopub.execute_input":"2025-03-30T04:10:35.617413Z","iopub.status.idle":"2025-03-30T04:10:36.18044Z","shell.execute_reply.started":"2025-03-30T04:10:35.617383Z","shell.execute_reply":"2025-03-30T04:10:36.179267Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from keras.callbacks import ReduceLROnPlateau, ModelCheckpoint, CSVLogger\n\ncallbacks = [\n    ReduceLROnPlateau(\n        monitor='loss',  # Change to monitor training loss (since val_loss may not be correctly passed)\n        factor=0.2,\n        patience=2,\n        min_lr=1e-5,\n        verbose=1\n    ),\n    ModelCheckpoint(\n        filepath='model_.{epoch:02d}-{loss:.6f}.weights.h5',  # Use training loss\n        monitor='loss',  # Monitor training loss\n        save_best_only=True,\n        save_weights_only=True,\n        verbose=1\n    ),\n    CSVLogger('training.log', separator=',', append=False)\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T05:08:15.003619Z","iopub.execute_input":"2025-03-30T05:08:15.003976Z","iopub.status.idle":"2025-03-30T05:08:15.008952Z","shell.execute_reply.started":"2025-03-30T05:08:15.003952Z","shell.execute_reply":"2025-03-30T05:08:15.008113Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 8. training model","metadata":{}},{"cell_type":"code","source":"K.clear_session()\nhistory =  model.fit(training_generator,\n                    epochs=35,\n                    steps_per_epoch=len(train_ids),\n                    callbacks= callbacks,\n                    validation_data = valid_generator\n                    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T04:48:46.478816Z","iopub.execute_input":"2025-03-30T04:48:46.479129Z","iopub.status.idle":"2025-03-30T04:50:48.239119Z","shell.execute_reply.started":"2025-03-30T04:48:46.479105Z","shell.execute_reply":"2025-03-30T04:50:48.237221Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"K.clear_session()\nhistory = model.fit(\n    training_generator,\n    epochs=35,\n    steps_per_epoch=steps_per_epoch,  # 125 步/epoch\n    callbacks=callbacks,\n    validation_data=valid_generator,\n    validation_steps=validation_steps\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-30T04:56:48.887316Z","iopub.execute_input":"2025-03-30T04:56:48.887619Z","iopub.status.idle":"2025-03-30T05:00:09.421418Z","shell.execute_reply.started":"2025-03-30T04:56:48.887595Z","shell.execute_reply":"2025-03-30T05:00:09.420117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}