{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\nimport glob\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\n\nsys.path.append('../input/stainnet/')\nfrom models import StainNet, ResnetGenerator","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-06-01T03:07:00.675488Z","iopub.execute_input":"2023-06-01T03:07:00.675919Z","iopub.status.idle":"2023-06-01T03:07:05.308900Z","shell.execute_reply.started":"2023-06-01T03:07:00.675882Z","shell.execute_reply":"2023-06-01T03:07:05.307962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# StainNet: make stain normalized image\n\nImage data of this competiton is composed of three dataset.　<br>\nAs you can see, he intensity of the staining of the slides in each dataset is different. <br>\n\nTherefore, color normalization of each dataset is necessary.  <br>\n\n In this notebook, we will introduce one of the methods of stain normalization, StainNet\n（https://www.frontiersin.org/articles/10.3389/fmed.2021.746307/full）. <br>\n\n## Please Upvote if you Find this Useful :)","metadata":{}},{"cell_type":"code","source":"tile_meta = pd.read_csv(\"/kaggle/input/hubmap-hacking-the-human-vasculature/tile_meta.csv\")\ntrain_imgs_root_path = \"/kaggle/input/hubmap-hacking-the-human-vasculature/train/\"\n\ndataset_1 = tile_meta.query(\"dataset==1\")[\"id\"].reset_index(drop=True)\ndataset_2 = tile_meta.query(\"dataset==2\")[\"id\"].reset_index(drop=True)\ndataset_3 = tile_meta.query(\"dataset==3\")[\"id\"].reset_index(drop=True)\ndatasets = [dataset_1, dataset_2, dataset_3]\n\nfig, axes = plt.subplots(1,3, figsize=(15, 10))\nfor j in range(3):\n    img = np.array(Image.open(train_imgs_root_path+datasets[j][0]+\".tif\"))\n    axes[j].imshow(img)\n    axes[j].tick_params(labelbottom=False, labelleft=False, labelright=False, labeltop=False)\n    axes[j].set_title(f\"Dataset_{j+1}\")","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-06-01T03:07:53.060582Z","iopub.execute_input":"2023-06-01T03:07:53.061021Z","iopub.status.idle":"2023-06-01T03:07:54.154638Z","shell.execute_reply.started":"2023-06-01T03:07:53.060968Z","shell.execute_reply":"2023-06-01T03:07:54.153779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# preparation for StainNet\nmodel_Net = StainNet().cuda()\nmodel_Net.load_state_dict(torch.load(\"../input/stainnet/checkpoints/aligned_histopathology_dataset/StainNet-Public_layer3_ch32.pth\"))\nmodel_Net.eval()\n\ndef norm(image):\n    image = np.array(image).astype(np.float32)\n    image = image.transpose((2, 0, 1))\n    image = ((image / 255) - 0.5) / 0.5\n    image=image[np.newaxis, ...]\n    image=torch.from_numpy(image)\n    return image\n\ndef un_norm(image):\n    image = image.cpu().detach().numpy()[0]\n    image = ((image * 0.5 + 0.5) * 255).astype(np.uint8).transpose((1,2,0))\n    return image\n\ndef stain_normalize(source, verbose=False):\n    with torch.no_grad():\n        img_net=model_Net(norm(source).cuda())\n        img_net=un_norm(img_net)\n        if verbose: plt.imshow(img_net); plt.show()\n        return img_net","metadata":{"execution":{"iopub.status.busy":"2023-06-01T02:35:22.673726Z","iopub.execute_input":"2023-06-01T02:35:22.676317Z","iopub.status.idle":"2023-06-01T02:35:25.992203Z","shell.execute_reply.started":"2023-06-01T02:35:22.676280Z","shell.execute_reply":"2023-06-01T02:35:25.991242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_imgs_root_path = \"/kaggle/input/hubmap-hacking-the-human-vasculature/train/\"\nsave_dir = \"/kaggle/working/stain_normalized_image\"\ntrain_imgs_path = glob.glob(train_imgs_root_path+\"*\")","metadata":{"execution":{"iopub.status.busy":"2023-06-01T02:35:25.994497Z","iopub.execute_input":"2023-06-01T02:35:25.995154Z","iopub.status.idle":"2023-06-01T02:35:26.295589Z","shell.execute_reply.started":"2023-06-01T02:35:25.995120Z","shell.execute_reply":"2023-06-01T02:35:26.294506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#create the directory to save normalized image\nos.mkdir(save_dir)\n\nfor img_path in tqdm(train_imgs_path):\n    \n    #load image\n    img = np.array(Image.open(img_path))\n    \n    #apply stain_normalization\n    img = stain_normalize(img)\n    \n    #save image\n    img = Image.fromarray(img)\n    img.save(save_dir+\"/\"+img_path.split(\"/\")[-1])\n    ","metadata":{"execution":{"iopub.status.busy":"2023-06-01T02:36:48.476263Z","iopub.execute_input":"2023-06-01T02:36:48.476636Z","iopub.status.idle":"2023-06-01T02:41:03.380019Z","shell.execute_reply.started":"2023-06-01T02:36:48.476604Z","shell.execute_reply":"2023-06-01T02:41:03.378439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"normalized_img_path = glob.glob(save_dir+\"/*\")\ntile_meta = pd.read_csv(\"/kaggle/input/hubmap-hacking-the-human-vasculature/tile_meta.csv\")\n\ndataset_1 = tile_meta.query(\"dataset==1\")[\"id\"].reset_index(drop=True)\ndataset_2 = tile_meta.query(\"dataset==2\")[\"id\"].reset_index(drop=True)\ndataset_3 = tile_meta.query(\"dataset==3\")[\"id\"].reset_index(drop=True)\ndatasets = [dataset_1, dataset_2, dataset_3]\n\nfig, axes = plt.subplots(2,3, figsize=(15, 10))\nfor j in range(3):\n    img = np.array(Image.open(train_imgs_root_path+datasets[j][1]+\".tif\"))\n    axes[0][j].imshow(img)\n    axes[0][j].tick_params(labelbottom=False, labelleft=False, labelright=False, labeltop=False)\n    axes[0][j].set_title(f\"Dataset_{j+1}: \\n Befor normalization\")\nfor j in range(3):\n    img = np.array(Image.open(save_dir+\"/\"+datasets[j][1]+\".tif\"))\n    axes[1][j].imshow(img)\n    axes[1][j].tick_params(labelbottom=False, labelleft=False, labelright=False, labeltop=False)\n    axes[1][j].set_title(f\"Dataset_{j+1}: \\n After normalization\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-01T02:54:37.035435Z","iopub.execute_input":"2023-06-01T02:54:37.036424Z","iopub.status.idle":"2023-06-01T02:54:38.683820Z","shell.execute_reply.started":"2023-06-01T02:54:37.036392Z","shell.execute_reply":"2023-06-01T02:54:38.682719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}