{"metadata":{"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":7326107,"sourceType":"datasetVersion","datasetId":4252213}],"dockerImageVersionId":30627,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import tempfile\nfrom pathlib import Path\nimport h5py\nimport tifffile\nimport numpy as np\nfrom fastai.vision.all import *\nimport dask","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:16.151075Z","iopub.execute_input":"2024-01-08T20:33:16.151985Z","iopub.status.idle":"2024-01-08T20:33:17.251755Z","shell.execute_reply.started":"2024-01-08T20:33:16.151954Z","shell.execute_reply":"2024-01-08T20:33:17.250803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FASTAI_ONECYCLE_N = 20\nFASTAI_FINETUNE_N = 20\nMLP_EPOCHS_N = 100\n\n#data contrast correction mode\nCORR_MODE=\"meanstd\" # Options are meanstd, minmax, div256\nCORR_SIGMA = 2.5 # only used when CORR_MODE=\"meanstd\"\n\nMODEL_PRETRAINED=models.resnet34\n\n#train_vol_files_root=\"/home/ypu66991/Desktop/Kaggle_KidneyBloodVessels/processed/traindata_v01_centre256_and_others/\"\ntrain_vol_files_root=\"/kaggle/input/traindata-v01/\"","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:17.253506Z","iopub.execute_input":"2024-01-08T20:33:17.254032Z","iopub.status.idle":"2024-01-08T20:33:17.259424Z","shell.execute_reply.started":"2024-01-08T20:33:17.254004Z","shell.execute_reply":"2024-01-08T20:33:17.258359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\ntry:\n    shutil.rmtree(\"fastai_train\")\nexcept:\n    pass\n\n# setup folders\ntemp_data_folder = Path(\"fastai_train/data\")\ntemp_labels_folder = Path(\"fastai_train/labels\")\n\ntemp_data_folder.mkdir(parents=True,exist_ok=True)\ntemp_labels_folder.mkdir(parents=True,exist_ok=True)","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:17.260482Z","iopub.execute_input":"2024-01-08T20:33:17.260783Z","iopub.status.idle":"2024-01-08T20:33:17.275958Z","shell.execute_reply.started":"2024-01-08T20:33:17.260759Z","shell.execute_reply":"2024-01-08T20:33:17.274988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(temp_data_folder)\nprint(temp_labels_folder)","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:17.277249Z","iopub.execute_input":"2024-01-08T20:33:17.277496Z","iopub.status.idle":"2024-01-08T20:33:17.291861Z","shell.execute_reply.started":"2024-01-08T20:33:17.277474Z","shell.execute_reply":"2024-01-08T20:33:17.290857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_vol_files = [\n#     \"/kaggle/input/traindata-v01/kidney1_256_roi1_clean.h5\",\n#     \"/kaggle/input/traindata-v01/kidney1_256_roi2_clean.h5\",\n#     \"/kaggle/input/traindata-v01/kidney1_256_roi3_clean.h5\",\n#     \"/kaggle/input/traindata-v01/kidney3_256_roi_clean.h5\",\n# ]\n\ntrain_vol_files = [\n    train_vol_files_root+\"kidney1_256_roi1_clean.h5\",\n    train_vol_files_root+\"kidney1_256_roi2_clean.h5\",\n    train_vol_files_root+\"kidney1_256_roi3_clean.h5\",\n    train_vol_files_root+\"kidney3_256_roi_clean.h5\",\n]\n","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:17.294703Z","iopub.execute_input":"2024-01-08T20:33:17.294975Z","iopub.status.idle":"2024-01-08T20:33:17.305374Z","shell.execute_reply.started":"2024-01-08T20:33:17.294951Z","shell.execute_reply":"2024-01-08T20:33:17.304715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def contrast_enhance_to_uint8(datavol0, sigma=2.5, mode=\"meanstd\"):\n    '''\n    mode:\n        meanstd, with default sigma of 2.5. Clips data to mean-sigma*std, meand+sigma*std, and adjusts to range 0-255 (uint8)\n    \n    '''\n    if mode==\"meanstd\":\n        std=np.std(datavol0)\n        mean=np.mean(datavol0)\n        data_corr = datavol0.astype(np.float32)-mean\n        datacorr0to1=np.clip(data_corr,-sigma*std, sigma*std)/(2*sigma*std)+0.5\n        data_corr_u8=(datacorr0to1*255).astype(np.uint8)\n        return data_corr_u8\n    \n    elif mode==\"minmax\":\n        min0=datavol0.min()\n        max0=datavol0.max()\n        data_corr_u8=((datavol0.astype(np.float32)-min0)/(max0-min0)*255).astype(np.uint8)\n        return data_corr_u8\n    \n    elif mode==\"div256\":\n        data_corr_u8 = (datavol0.astype(np.float32)/256).astype(np.uint8)\n        return data_corr_u8\n    else:\n        return datavol0.astype(np.uint8) # Attention, this should be avoided\n        ","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:17.306648Z","iopub.execute_input":"2024-01-08T20:33:17.307292Z","iopub.status.idle":"2024-01-08T20:33:17.321273Z","shell.execute_reply.started":"2024-01-08T20:33:17.307256Z","shell.execute_reply":"2024-01-08T20:33:17.320565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import functools\nmy_corr = functools.partial(contrast_enhance_to_uint8, sigma=CORR_SIGMA, mode=CORR_MODE)\n# please make sure to use my_corr at inference too","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:17.322766Z","iopub.execute_input":"2024-01-08T20:33:17.323107Z","iopub.status.idle":"2024-01-08T20:33:17.335403Z","shell.execute_reply.started":"2024-01-08T20:33:17.323076Z","shell.execute_reply":"2024-01-08T20:33:17.334607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# generated files for training\n\nicount=0\naxis=['z','y','x']\n\nfor f0 in train_vol_files:\n    print(\"file:\", f0)\n    #load data and labels\n    with h5py.File(f0) as f:\n        datavol=np.array(f['raw'])\n        labelvol=np.array(f['label'])\n    \n    # correct volume and convert to uint8\n    #std=np.std(datavol)\n    #mean=np.mean(datavol)\n    #sigma=2.5\n    #data_corr = datavol.astype(np.float32)-mean\n    #datacorr0to1=np.clip(data_corr,-sigma*std, sigma*std)/(2*sigma*std)+0.5\n    #data_corr_u8=(datacorr0to1*255).astype(np.uint8)\n    #data_corr_u8 = ((datavol.astype(np.float32)-datavol.min())/ (datavol.max()-datavol.min())*255).astype(np.uint8)\n    #data_corr_u8= datavol\n    data_corr_u8 = my_corr(datavol)\n    \n    #Make labels mask=1\n    #labelvol_u8 = np.where(labelvol>0, 1, 0).astype(np.uint8)\n    labelvol_u8= labelvol\n    \n    for ax0 in axis:\n        print(f\"axis:{ax0}\")\n        # grab slice depending in axis\n        if ax0=='y':\n            datavol_rot=np.transpose(data_corr_u8, axes=(1,2,0))\n            datalabel_rot= np.transpose(labelvol_u8, axes=(1,2,0))\n        elif ax0=='x':\n            datavol_rot=np.transpose(data_corr_u8, axes=(2,0,1))\n            datalabel_rot=np.transpose(labelvol_u8, axes=(2,0,1))\n        else:\n            datavol_rot=data_corr_u8\n            datalabel_rot=labelvol_u8\n        \n        for islice in range(datavol_rot.shape[0]):\n            slice_d=datavol_rot[islice,:,:]\n            slice_l=datalabel_rot[islice,:,:]\n              \n            #save slice(s) to disk, no corrections\n            fn_save = f\"{icount:04d}.tif\"\n            \n            #slice data save\n            tifffile.imwrite( Path(temp_data_folder)/fn_save , slice_d)\n            tifffile.imwrite( Path(temp_labels_folder)/fn_save , slice_l)\n            \n            icount+=1\n            \n            #if icount>=10:\n            #    break","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:17.336512Z","iopub.execute_input":"2024-01-08T20:33:17.336906Z","iopub.status.idle":"2024-01-08T20:33:25.638996Z","shell.execute_reply.started":"2024-01-08T20:33:17.336856Z","shell.execute_reply":"2024-01-08T20:33:25.638086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!ls {temp_data_folder.name}","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:25.640182Z","iopub.execute_input":"2024-01-08T20:33:25.640469Z","iopub.status.idle":"2024-01-08T20:33:25.644514Z","shell.execute_reply.started":"2024-01-08T20:33:25.640444Z","shell.execute_reply":"2024-01-08T20:33:25.643594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setup fastai","metadata":{}},{"cell_type":"code","source":"fnames = get_image_files(temp_data_folder)\nfnames","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:25.6457Z","iopub.execute_input":"2024-01-08T20:33:25.645985Z","iopub.status.idle":"2024-01-08T20:33:25.683126Z","shell.execute_reply.started":"2024-01-08T20:33:25.645961Z","shell.execute_reply":"2024-01-08T20:33:25.682335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# function that converts data filename to label\ndef get_lbl_from_fn(fn):\n    return Path(temp_labels_folder)/ f\"{fn.stem}.tif\"","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:25.684334Z","iopub.execute_input":"2024-01-08T20:33:25.685091Z","iopub.status.idle":"2024-01-08T20:33:25.689487Z","shell.execute_reply.started":"2024-01-08T20:33:25.685058Z","shell.execute_reply":"2024-01-08T20:33:25.688617Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test\nget_lbl_from_fn(fnames[0]).exists()","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:25.691058Z","iopub.execute_input":"2024-01-08T20:33:25.691401Z","iopub.status.idle":"2024-01-08T20:33:25.707506Z","shell.execute_reply.started":"2024-01-08T20:33:25.691369Z","shell.execute_reply":"2024-01-08T20:33:25.706704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegmentationDataLoadersAndBlockBW(DataLoaders):\n    # Modified version of SegmentationDataLoaders for greyscale images\n    @classmethod\n    @delegates(DataLoaders.from_dblock)\n    def from_label_func(cls, path, fnames, label_func, valid_pct=0.2, seed=None, codes=None, item_tfms=None, batch_tfms=None, \n                        **kwargs):\n        \"Create from list of `fnames` in `path`s with `label_func`.\"\n        dblock = DataBlock(blocks=(ImageBlock(cls=PILImageBW), MaskBlock(codes=codes)),\n                           splitter=RandomSplitter(valid_pct, seed=seed),\n                           get_y=label_func,\n                           item_tfms=item_tfms,\n                           batch_tfms=batch_tfms)\n        #batch_tfms=[IntToFloatTensor(div=2**16-1), *aug_transforms()])\n        #mean,std = [0.5]*3,[0.5]*3\n        #mean,std = broadcast_vec(1, 4, mean, std)\n        #dblock = DataBlock(blocks=(ImageBlock(cls=PILImageBW), MaskBlock(codes=codes)),\n        #           splitter=RandomSplitter(valid_pct, seed=seed),\n        #           get_y=label_func,\n        #           item_tfms=item_tfms,\n        #       #batch_tfms=[IntToFloatTensor(div=65536,div_mask=1), *aug_transforms()])\n        #           batch_tfms=[IntToFloatTensor(div=65536,div_mask=1), Normalize(), *aug_transforms()])\n        res = cls.from_dblock(dblock, fnames, path=path, **kwargs)\n        return res , dblock","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:25.708511Z","iopub.execute_input":"2024-01-08T20:33:25.70877Z","iopub.status.idle":"2024-01-08T20:33:25.719897Z","shell.execute_reply.started":"2024-01-08T20:33:25.708748Z","shell.execute_reply":"2024-01-08T20:33:25.719014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path_model_root= Path(\"multiaxis_models\")","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:25.724794Z","iopub.execute_input":"2024-01-08T20:33:25.725533Z","iopub.status.idle":"2024-01-08T20:33:25.730395Z","shell.execute_reply.started":"2024-01-08T20:33:25.7255Z","shell.execute_reply":"2024-01-08T20:33:25.729536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imdl1, bl1 = SegmentationDataLoadersAndBlockBW.from_label_func(path_model_root,bs=4, fnames=fnames ,label_func= get_lbl_from_fn, codes=[\"bkg\",\"vessel\"])","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:25.73138Z","iopub.execute_input":"2024-01-08T20:33:25.731651Z","iopub.status.idle":"2024-01-08T20:33:26.340631Z","shell.execute_reply.started":"2024-01-08T20:33:25.731627Z","shell.execute_reply":"2024-01-08T20:33:26.339523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imdl1.show_batch()","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:26.341833Z","iopub.execute_input":"2024-01-08T20:33:26.342111Z","iopub.status.idle":"2024-01-08T20:33:26.904294Z","shell.execute_reply.started":"2024-01-08T20:33:26.342087Z","shell.execute_reply":"2024-01-08T20:33:26.903272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"OK","metadata":{}},{"cell_type":"code","source":"#learn = unet_learner(imdl1, models.resnet34, n_in=1)\n#learn = unet_learner(imdl1, models.resnet50, n_in=1)\nlearn = unet_learner(imdl1, MODEL_PRETRAINED, n_in=1)","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:26.905951Z","iopub.execute_input":"2024-01-08T20:33:26.906254Z","iopub.status.idle":"2024-01-08T20:33:29.452667Z","shell.execute_reply.started":"2024-01-08T20:33:26.906226Z","shell.execute_reply":"2024-01-08T20:33:29.451784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn.summary()","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:29.454048Z","iopub.execute_input":"2024-01-08T20:33:29.454436Z","iopub.status.idle":"2024-01-08T20:33:30.93676Z","shell.execute_reply.started":"2024-01-08T20:33:29.454399Z","shell.execute_reply":"2024-01-08T20:33:30.935662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn.lr_find()","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:30.938008Z","iopub.execute_input":"2024-01-08T20:33:30.938312Z","iopub.status.idle":"2024-01-08T20:33:46.598444Z","shell.execute_reply.started":"2024-01-08T20:33:30.938286Z","shell.execute_reply":"2024-01-08T20:33:46.597408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Fitting","metadata":{}},{"cell_type":"code","source":"learn.fit_one_cycle(FASTAI_ONECYCLE_N) #Add more cycles if needed","metadata":{"execution":{"iopub.status.busy":"2024-01-08T20:33:46.599886Z","iopub.execute_input":"2024-01-08T20:33:46.600172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#import dill\n#learn.export(\"kidney_fastai_multiaxis_chkp01.dpkl\", pickle_module=dill)\n#Load it using load_lerner(\"<name>.dpkl\", pickle_module=dill)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn.unfreeze()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn.fine_tune(FASTAI_FINETUNE_N) #Add more if needed","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import dill\nimport datetime\n\nDATE= datetime.date.today()\n\n#learn.export(f\"{DATE}_kidney_multiaxis_fastai_corr_{CORR_MODE}_sigma_{CORR_SIGMA}.dpkl\", pickle_module=dill)\n#Load it using load_lerner(\"<name>.dpkl\", pickle_module=dill, cpu=False)\n\nlearn.export(f\"{DATE}_kidney_multiaxis_fastai_corr_{CORR_MODE}_sigma_{CORR_SIGMA}_resnet50.dpkl\", pickle_module=dill)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Cleanup after fastai unet training","metadata":{}},{"cell_type":"code","source":"# Cleanup temporary files used for training\nimport shutil\ntry:\n    shutil.rmtree(\"fastai_train\")\nexcept:\n    pass","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# pre-MLP, Run multiaxis predictions on each volume","metadata":{}},{"cell_type":"code","source":"import shutil\ntry:\n    shutil.rmtree(\"fastai_preds\")\nexcept:\n    pass\ntemp_pred_folder = Path(\"fastai_preds\")\ntemp_pred_folder.mkdir(parents=True,exist_ok=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"icount=0\naxis=['z','y','x']\n\nsaczyx =[] #list that is later stacked, set,predaxis, z,y,x,class\nlabels_szyx=[]\n\nfor f0 in train_vol_files:\n    print(\"file:\", f0)\n    #load data and labels\n    with h5py.File(f0) as f:\n        datavol=np.array(f['raw'])\n        labelvol=np.array(f['label'])\n    \n    # correct volume and convert to uint8\n    #std=np.std(datavol)\n    #mean=np.mean(datavol)\n    #sigma=2.5\n    #data_corr = datavol.astype(np.float32)-mean\n    #datacorr0to1=np.clip(data_corr,-sigma*std, sigma*std)/(2*sigma*std)+0.5\n    #data_corr_u8=(datacorr0to1*255).astype(np.uint8)\n    data_corr_u8=my_corr(datavol)\n    \n    czyx =[] #list that is later stacked, predaxis, z,class, y,x\n    labels_szyx.append(labelvol)\n    \n    for ax0 in axis:\n        print(f\"axis:{ax0}\")\n        # grab slice depending in axis\n        if ax0=='y':\n            datavol_rot=np.transpose(data_corr_u8, axes=(1,2,0)) # becomes y-xz\n        elif ax0=='x':\n            datavol_rot=np.transpose(data_corr_u8, axes=(2,0,1)) # becomes x-zy\n        else:\n            datavol_rot=data_corr_u8\n        \n        vol_pred_probs_slices=[]\n        #slice by slice\n        for islice in range(datavol_rot.shape[0]):\n            slice_d=datavol_rot[islice,:,:]\n            \n            # Run prediction\n            # learn_pred = learn.predict(slice_d)\n            with learn.no_bar(), learn.no_logging():\n                learn_pred = learn.predict(slice_d)\n            # Collect probabilities\n            probs = np.array(learn_pred[2])\n            vol_pred_probs_slices.append(probs)\n        \n        vol_pred_probs = np.array(vol_pred_probs_slices) #stack all slices\n        print(f\"vol_pred_probs.shape: {vol_pred_probs.shape}\")\n        \n        #Correct for axis\n        if ax0=='y':\n            #yaxz to azyx\n            vol_pred_probs_corr=np.transpose(vol_pred_probs, axes=(1,3,0,2))\n        elif ax0=='x':\n            #xazy to azyx\n            vol_pred_probs_corr=np.transpose(vol_pred_probs, axes=(1,2,3,0))\n        else:#zayx to azyx\n            vol_pred_probs_corr=np.transpose(vol_pred_probs, axes=(1,0,2,3))\n              \n        #Store result in volprobs_all_np4d\n        czyx.append(vol_pred_probs_corr)\n    \n    #Completed axis, now stack\n    aczyx_np = np.array(czyx)\n    saczyx.append(aczyx_np)\n    \n#Completed set, now stack\nsaczyx_np = np.array(saczyx)\nlabels_szyx_np = np.array(labels_szyx)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"saczyx_np.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#check\nimport matplotlib.pyplot as plt\n\niset=0\nfig,ax= plt.subplots(1,4, figsize=(10,5))\n\nax[0].imshow(saczyx_np[0,0,1,128,:,:])\nax[0].set_axis_off()\nax[1].imshow(saczyx_np[0,1,1,128,:,:])\nax[1].set_axis_off()\nax[2].imshow(saczyx_np[0,2,1,128,:,:])\nax[2].set_axis_off()\nax[3].imshow(labels_szyx_np[0,128,:,:], cmap='gray')\nax[3].set_axis_off()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"OK","metadata":{}},{"cell_type":"code","source":"print(f\"saczyx_np.shape:{saczyx_np.shape}\")\nprint(f\"labels_szyx_np.shape:{labels_szyx_np.shape}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Manipulate data and labels to make it ready for MLP classification.\n\nNeed to flatten on set and along zyx axis. Then later put channel as last dim. axis and set can be flatten to become batch","metadata":{}},{"cell_type":"code","source":"test0 = np.transpose(saczyx_np, axes=[1,2,0,3,4,5])\ntest0.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# flatten set,zyx\ntest1 = np.reshape(test0, (*test0.shape[:2], np.prod(test0.shape[2:]) ) )\ntest1.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test2 = np.transpose(test1, axes=[2,0,1])\ntest2.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#flatten axis, nprobs\ntest3 = np.reshape(test2, (test2.shape[0], np.prod(test2.shape[1:] ) ) )\ntest3.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_flat_for_mlp = test3\ndata_flat_for_mlp.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"flatten labels in a similar way","metadata":{}},{"cell_type":"code","source":"labels_szyx_np.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_test0 = np.reshape(labels_szyx_np, (np.prod(labels_szyx_np.shape),1))\nlabel_test0.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setup pytorch MLP classifier","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import TensorDataset, DataLoader   ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n#device\n\n#device=torch.device('cpu')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define the MLP model\nclass MLPClassifier(nn.Module):\n    def __init__(self, input_size, hidden_size1, hidden_size2, output_size):\n        super(MLPClassifier, self).__init__()\n        self.fc1 = nn.Linear(input_size, hidden_size1)\n        self.fc2 = nn.Linear(hidden_size1, hidden_size2)\n        self.fc3 = nn.Linear(hidden_size2, output_size)\n        self.tanh = nn.Tanh()\n\n    def forward(self, x):\n        x = self.tanh(self.fc1(x))\n        x = self.tanh(self.fc2(x))\n        x = self.fc3(x)\n        #x = self.tanh(self.fc3(x))\n        return x","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set the input size, hidden layer sizes, and output size\ninput_size = 6\nhidden_size1 = 10\nhidden_size2 = 10\noutput_size = 1  # Binary classification, so output size is 1","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create the model\nmodel_MLP_fusion = MLPClassifier(input_size, hidden_size1, hidden_size2, output_size)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_MLP_fusion.to(device)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# define loss function and optimizer\ncriterion = nn.BCEWithLogitsLoss()\n#criterion = nn.BCELoss()\n#optimizer = optim.SGD(model_MLP_fusion.parameters(), lr=0.00001)\noptimizer = optim.Adam(model_MLP_fusion.parameters(), lr=0.001)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train= torch.from_numpy(data_flat_for_mlp).to(device)\ny_train= torch.from_numpy(label_test0.astype(np.float32)).to(device)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Set the size of the subset you want to use for training\nsubset_size = 2**18  # Adjust this based on your preferences\n\n# Create a random subset of the indices\nsubset_indices = torch.randperm(len(X_train))[:subset_size]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Use the subset indices to create a TensorDataset\nsubset_dataset = TensorDataset(X_train[subset_indices], y_train[subset_indices])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create a DataLoader with the dataset and set shuffle=True\n# Combine X_train and y_train into a TensorDataset\ntrain_dataset = TensorDataset(X_train, y_train)\n\nbatch_size = 256\n#train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\ntrain_loader = DataLoader(subset_dataset, batch_size=batch_size, shuffle=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Test\n#mytrainlist = list(train_loader) #shuffling takes too long","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training loop\nnum_epochs = MLP_EPOCHS_N\nlosses=[]\nfor epoch in range(num_epochs):\n    for batch_X, batch_y in train_loader:\n         # Move batch data to the GPU if available\n        #batch_X, batch_y = batch_X.to(device), batch_y.to(device)\n        \n        # Forward pass\n        outputs = model_MLP_fusion(batch_X)\n        loss = criterion(outputs, batch_y)\n\n        # Backward pass and optimization\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n    #if (epoch + 1) % 10 == 0:\n    print(f'Epoch [{epoch + 1}/{num_epochs}], Loss: {loss.item():.8f}')\n    losses.append(float(loss))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save the trained model\ntorch.save(model_MLP_fusion.state_dict(), f'multiaxis_models/{DATE}_kidney_fastai_MLPfusion_{CORR_MODE}_sigma_{CORR_SIGMA}.pth')\n\n# To load the model later\n#loaded_model = MLPClassifier(input_size, hidden_size1, hidden_size2, output_size)\n#loaded_model.load_state_dict(torch.load('trained_model.pth'))\n#loaded_model.eval()  # Set the model to evaluation mode if needed","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict with validation data and surface dice metric","metadata":{}},{"cell_type":"code","source":"def NN1_predict(datavol0, fai_learner0):\n    '''\n    Returns:\n        a stack of volumes with predictions along the different axis/planes as a np array\n        in the order axis,channel, z,y,x\n    '''\n    \n    axis=['z','y','x']\n    \n    #Adjust data contrast. by mean and stdev\n    \n    # correct volume and convert to uint8\n    data_corr_u8 = contrast_enhance_to_uint8(datavol0)\n    \n    aczyx =[] #list that is later stacked: class, z,class, y,x\n    \n    for ax0 in axis:\n        print(f\"axis:{ax0}\")\n        # grab slice depending in axis\n        if ax0=='y':\n            datavol_rot=np.transpose(data_corr_u8, axes=(1,2,0)) # becomes y-xz\n        elif ax0=='x':\n            datavol_rot=np.transpose(data_corr_u8, axes=(2,0,1)) # becomes x-zy\n        else:\n            datavol_rot=data_corr_u8\n        \n        vol_pred_probs_slices=[]\n        #slice by slice\n        for islice in range(datavol_rot.shape[0]):\n            slice_d=datavol_rot[islice,:,:]\n            \n            # Run prediction\n            #learn_pred = fai_learner0.predict(slice_d)\n            # To run prediction without printing the bar\n            with fai_learner0.no_bar(), fai_learner0.no_logging():\n                learn_pred = fai_learner0.predict(slice_d)\n            \n            # Collect probabilities\n            probs = np.array(learn_pred[2])\n            vol_pred_probs_slices.append(probs)\n        \n        vol_pred_probs = np.array(vol_pred_probs_slices) #stack all slices\n        print(f\"vol_pred_probs.shape: {vol_pred_probs.shape}\")\n        \n        #Correct for axis\n        if ax0=='y':\n            #yaxz to azyx\n            vol_pred_probs_corr=np.transpose(vol_pred_probs, axes=(1,3,0,2))\n        elif ax0=='x':\n            #xazy to azyx\n            vol_pred_probs_corr=np.transpose(vol_pred_probs, axes=(1,2,3,0))\n        else:#zayx to azyx\n            vol_pred_probs_corr=np.transpose(vol_pred_probs, axes=(1,0,2,3))\n              \n        #Store result in volprobs_all_np4d\n        aczyx.append(vol_pred_probs_corr)\n    \n    #Completed axis, now settle stack to np\n    aczyx_np = np.array(aczyx)\n    \n    return aczyx_np","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def NN2_predict(aczyx0, mlp_fusion0, device0):\n    mlp_fusion0.eval()\n    \n    preds_multiplane = aczyx0\n    # expect data in format planeaxis, probclass, z,y,x\n    assert preds_multiplane.ndim==5\n    \n    #Flatten along zyx\n    temp0 = np.reshape(preds_multiplane, (*preds_multiplane.shape[:2], np.prod(preds_multiplane.shape[2:])))\n    #print(f\"temp0.shape: {temp0.shape}\")\n    \n    #transpose\n    temp1 = np.transpose(temp0, axes=[2,0,1])\n    #flatten along the multipredaxis and probs\n    temp2 = np.reshape(temp1, (temp1.shape[0], np.prod(temp1.shape[1:])))\n    \n    inp_x_tensor = torch.from_numpy(temp2).to(device0) # maybe I should wrap this in a class to ensure device is the same as mlp\n    # Run prediction with this data\n    res = mlp_fusion0(inp_x_tensor)\n    #print(f\"res.shape:{res.shape}, res.max:{res.max()}, res.min:{res.min()}\")\n    \n    #Get the label\n    labels_flat_np = res.cpu().detach().numpy()\n    labels_bin = (labels_flat_np>0).astype(np.uint8)\n    #Unflatten\n    labels = np.reshape(labels_bin, preds_multiplane.shape[2:] )\n    \n    return labels","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mlp_device=device","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_full(datavol):\n    global learn\n    global model_MLP_fusion\n    global mlp_device\n    \n    #Data can be dask, so convert to numpy\n    if dask.is_dask_collection(datavol):\n        datavol=datavol.compute()\n    \n    nn1_res = NN1_predict(datavol,learn)\n    #note that data is contrast corrected and converted to uint8 within NN1_predict\n    \n    nn2_res = NN2_predict(nn1_res, model_MLP_fusion, mlp_device)\n    \n    return nn2_res","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#with h5py.File(\"/home/ypu66991/Desktop/Kaggle_KidneyBloodVessels/processed/traindata_v01_centre256_and_others/kidney1_512_roi.h5\") as f:\nwith h5py.File(train_vol_files_root+\"kidney1_512_roi.h5\") as f:\n    data_for_pred=np.array(f['raw'])\n    labels_gnd = np.array(f['label'])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Run the full prediction here","metadata":{}},{"cell_type":"code","source":"pred_vol=predict_full(data_for_pred)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Metrics","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append(\"..\")\nimport metric","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask_gt = labels_gnd>0\nmask_pred = pred_vol>0","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"surf_dist = metric.compute_surface_distances(mask_gt,mask_pred, spacing_mm=(50e-3, 50e-3, 50e-3))\nsurf_dice_tol0 = metric.compute_surface_dice_at_tolerance(surf_dist,tolerance_mm=0)\nsurf_dice_tol0","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"iz=200\nimport matplotlib.pyplot as plt\nfig,ax= plt.subplots(1,3, figsize=(15,10))\nax[0].imshow(data_for_pred[iz,:,:])\nax[0].set_axis_off()\nax[1].imshow(labels_gnd[iz,:,:])\nax[1].set_axis_off()\nax[2].imshow(pred_vol[iz,:,:])\nax[2].set_axis_off()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}