{"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 numpy as np\nimport pandas as pd\nimport os\nfrom glob import glob\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport pydicom\nimport matplotlib.pyplot as plt\nimport cv2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-15T03:19:38.129866Z","iopub.execute_input":"2023-01-15T03:19:38.130648Z","iopub.status.idle":"2023-01-15T03:19:38.721912Z","shell.execute_reply.started":"2023-01-15T03:19:38.130555Z","shell.execute_reply":"2023-01-15T03:19:38.720821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_train_labels.csv')","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:38.724231Z","iopub.execute_input":"2023-01-15T03:19:38.725089Z","iopub.status.idle":"2023-01-15T03:19:38.762010Z","shell.execute_reply.started":"2023-01-15T03:19:38.725050Z","shell.execute_reply":"2023-01-15T03:19:38.761025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head(10)","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:38.763565Z","iopub.execute_input":"2023-01-15T03:19:38.763990Z","iopub.status.idle":"2023-01-15T03:19:38.785105Z","shell.execute_reply.started":"2023-01-15T03:19:38.763928Z","shell.execute_reply":"2023-01-15T03:19:38.783979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(df)","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:38.788537Z","iopub.execute_input":"2023-01-15T03:19:38.788920Z","iopub.status.idle":"2023-01-15T03:19:38.796289Z","shell.execute_reply.started":"2023-01-15T03:19:38.788884Z","shell.execute_reply":"2023-01-15T03:19:38.795142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataLoader(torch.utils.data.Dataset):\n    def __init__(self, df, bs):\n        self.df = df\n        self.bs = bs\n    def __len__(self):\n        return len(self.df) // self.bs\n    def __getitem__(self, _):\n        batch_imgs = []\n        batch_msks = []\n        for b in range(self.bs):\n            df_shuffled = self.df.sample(frac=1).reset_index(drop=True)\n            ptn = df_shuffled.patientId.tolist()[0]\n            target = float(df_shuffled.Target.tolist()[0])\n            if target>0:\n                x = int(df_shuffled.x.tolist()[0])\n                y = int(df_shuffled.y.tolist()[0])\n                width = int(df_shuffled.width.tolist()[0])\n                height= int(df_shuffled.height.tolist()[0])\n            else:\n                x, y, width, height = 0, 0, 0, 0\n            dcm = pydicom.dcmread(\n                os.path.join(\"/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_train_images\", \n                             ptn + '.dcm')\n            )\n            npy = dcm.pixel_array # + float(dcm[0x281052])\n            msk = np.zeros(npy.shape)\n            msk[x:x+width, y:y+width] = 1\n            npy = cv2.resize(npy, (512, 512), cv2.INTER_LINEAR)\n            msk = cv2.resize(msk, (512, 512), cv2.INTER_NEAREST)\n            batch_msks.append(msk)\n            batch_imgs.append(self._StochasticWindowing(npy))\n        imgs = torch.from_numpy(np.array(batch_imgs)).squeeze().unsqueeze(1)\n        msks = torch.from_numpy(np.array(batch_msks)).squeeze().unsqueeze(1)\n        \n        return imgs.float(), msks.float()\n    def _StochasticWindowing(self, img):\n        per = np.percentile(img, 95)\n        img = img / per\n        img[img>1] = 1\n        mean = img.mean()\n        std = img.std()\n        return (img - mean) / np.abs(np.random.normal()/2 + 1.96 * std)","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:38.798151Z","iopub.execute_input":"2023-01-15T03:19:38.798881Z","iopub.status.idle":"2023-01-15T03:19:38.815013Z","shell.execute_reply.started":"2023-01-15T03:19:38.798843Z","shell.execute_reply":"2023-01-15T03:19:38.813803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    \"\"\"(convolution => [BN] => ReLU) * 2\"\"\"\n\n    def __init__(self, in_channels, out_channels, mid_channels=None):\n        super().__init__()\n        if not mid_channels:\n            mid_channels = out_channels\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\n\nclass Down(nn.Module):\n    \"\"\"Downscaling with maxpool then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super().__init__()\n\n        # if bilinear, use the normal convolutions to reduce the number of channels\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n            self.conv = DoubleConv(in_channels, out_channels, in_channels // 2)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)\n            self.conv = DoubleConv(in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        # input is CHW\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,\n                        diffY // 2, diffY - diffY // 2])\n        # if you have padding issues, see\n        # https://github.com/HaiyongJiang/U-Net-Pytorch-Unstructured-Buggy/commit/0e854509c2cea854e247a9c615f175f76fbb2e3a\n        # https://github.com/xiaopeng-liao/Pytorch-UNet/commit/8ebac70e633bac59fc22bb5195e513d5832fb3bd\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\n\nclass OutConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(OutConv, self).__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:38.819114Z","iopub.execute_input":"2023-01-15T03:19:38.819442Z","iopub.status.idle":"2023-01-15T03:19:38.835932Z","shell.execute_reply.started":"2023-01-15T03:19:38.819413Z","shell.execute_reply":"2023-01-15T03:19:38.834852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, n_channels, n_classes, bilinear=False):\n        super(UNet, self).__init__()\n        self.n_channels = n_channels\n        self.n_classes = n_classes\n        self.bilinear = bilinear\n\n        self.inc = (DoubleConv(n_channels, 64))\n        self.down1 = (Down(64, 128))\n        self.down2 = (Down(128, 256))\n        self.down3 = (Down(256, 512))\n        factor = 2 if bilinear else 1\n        self.down4 = (Down(512, 1024 // factor))\n        self.up1 = (Up(1024, 512 // factor, bilinear))\n        self.up2 = (Up(512, 256 // factor, bilinear))\n        self.up3 = (Up(256, 128 // factor, bilinear))\n        self.up4 = (Up(128, 64, bilinear))\n        self.outc = (OutConv(64, n_classes))\n\n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        x = self.up1(x5, x4)\n        x = self.up2(x, x3)\n        x = self.up3(x, x2)\n        x = self.up4(x, x1)\n        logits = self.outc(x)\n        return logits\n\n    def use_checkpointing(self):\n        self.inc = torch.utils.checkpoint(self.inc)\n        self.down1 = torch.utils.checkpoint(self.down1)\n        self.down2 = torch.utils.checkpoint(self.down2)\n        self.down3 = torch.utils.checkpoint(self.down3)\n        self.down4 = torch.utils.checkpoint(self.down4)\n        self.up1 = torch.utils.checkpoint(self.up1)\n        self.up2 = torch.utils.checkpoint(self.up2)\n        self.up3 = torch.utils.checkpoint(self.up3)\n        self.up4 = torch.utils.checkpoint(self.up4)\n        self.outc = torch.utils.checkpoint(self.outc)","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:38.837362Z","iopub.execute_input":"2023-01-15T03:19:38.838621Z","iopub.status.idle":"2023-01-15T03:19:38.851554Z","shell.execute_reply.started":"2023-01-15T03:19:38.838570Z","shell.execute_reply":"2023-01-15T03:19:38.850606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net = UNet(1, 1).float()","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:38.852956Z","iopub.execute_input":"2023-01-15T03:19:38.853444Z","iopub.status.idle":"2023-01-15T03:19:39.159617Z","shell.execute_reply.started":"2023-01-15T03:19:38.853406Z","shell.execute_reply":"2023-01-15T03:19:39.158579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dloader = DataLoader(df, 8)\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nopt = torch.optim.Adam(net.parameters())\nnet = net.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:39.161370Z","iopub.execute_input":"2023-01-15T03:19:39.161787Z","iopub.status.idle":"2023-01-15T03:19:40.995937Z","shell.execute_reply.started":"2023-01-15T03:19:39.161744Z","shell.execute_reply":"2023-01-15T03:19:40.994898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n    def __init__(self):\n        super(DiceLoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        \n        inputs = F.sigmoid(inputs) # sigmoid를 통과한 출력이면 주석처리\n        \n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()                            \n        dice = (2.*intersection + smooth) / (inputs.sum() + targets.sum() + smooth)  \n        \n        return 1 - dice ","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:40.997271Z","iopub.execute_input":"2023-01-15T03:19:40.997637Z","iopub.status.idle":"2023-01-15T03:19:41.007134Z","shell.execute_reply.started":"2023-01-15T03:19:40.997598Z","shell.execute_reply":"2023-01-15T03:19:41.005901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LossFx = DiceLoss()\ntotLoss = []\nfor i in range(len(dloader)):\n    img, msk = dloader[i]\n    out = net(img.to(device))\n    loss = LossFx(out, msk.to(device))\n    opt.zero_grad()\n    loss.backward()\n    opt.step()\n    totLoss.append(loss.item())\n    print(f\"Iteration {i}/{len(dloader)} | DICE Loss {np.mean(totLoss)} | DICE {1-np.mean(totLoss)}\")","metadata":{"execution":{"iopub.status.busy":"2023-01-15T03:19:41.008745Z","iopub.execute_input":"2023-01-15T03:19:41.009269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}