{"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":"!pip install natsort\nimport numpy as np\nimport os\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom natsort import natsorted\nimport pydicom\nfrom glob import glob","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-14T16:03:49.268072Z","iopub.execute_input":"2023-01-14T16:03:49.269109Z","iopub.status.idle":"2023-01-14T16:03:59.747131Z","shell.execute_reply.started":"2023-01-14T16:03:49.269042Z","shell.execute_reply":"2023-01-14T16:03:59.746125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"basepath = '/kaggle/input/osic-pulmonary-fibrosis-progression'\ntrain_csv = pd.read_csv(basepath + '/train.csv')","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:03:59.753311Z","iopub.execute_input":"2023-01-14T16:03:59.755628Z","iopub.status.idle":"2023-01-14T16:03:59.766588Z","shell.execute_reply.started":"2023-01-14T16:03:59.755564Z","shell.execute_reply":"2023-01-14T16:03:59.765715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv.head(10)","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:03:59.767864Z","iopub.execute_input":"2023-01-14T16:03:59.769184Z","iopub.status.idle":"2023-01-14T16:03:59.794817Z","shell.execute_reply.started":"2023-01-14T16:03:59.768613Z","shell.execute_reply":"2023-01-14T16:03:59.793897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nfrom functools import partial\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\ndef get_inplanes():\n    return [64, 128, 256, 512]\n\n\ndef conv3x3x3(in_planes, out_planes, stride=1):\n    return nn.Conv3d(in_planes,\n                     out_planes,\n                     kernel_size=3,\n                     stride=stride,\n                     padding=1,\n                     bias=False)\n\n\ndef conv1x1x1(in_planes, out_planes, stride=1):\n    return nn.Conv3d(in_planes,\n                     out_planes,\n                     kernel_size=1,\n                     stride=stride,\n                     bias=False)\n\n\nclass BasicBlock(nn.Module):\n    expansion = 1\n\n    def __init__(self, in_planes, planes, stride=1, downsample=None):\n        super().__init__()\n\n        self.conv1 = conv3x3x3(in_planes, planes, stride)\n        self.bn1 = nn.BatchNorm3d(planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = conv3x3x3(planes, planes)\n        self.bn2 = nn.BatchNorm3d(planes)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out += residual\n        out = self.relu(out)\n\n        return out\n\n\nclass Bottleneck(nn.Module):\n    expansion = 4\n\n    def __init__(self, in_planes, planes, stride=1, downsample=None):\n        super().__init__()\n\n        self.conv1 = conv1x1x1(in_planes, planes)\n        self.bn1 = nn.BatchNorm3d(planes)\n        self.conv2 = conv3x3x3(planes, planes, stride)\n        self.bn2 = nn.BatchNorm3d(planes)\n        self.conv3 = conv1x1x1(planes, planes * self.expansion)\n        self.bn3 = nn.BatchNorm3d(planes * self.expansion)\n        self.relu = nn.ReLU(inplace=True)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n\n        out = self.conv3(out)\n        out = self.bn3(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out += residual\n        out = self.relu(out)\n\n        return out\n\n\nclass ResNet(nn.Module):\n\n    def __init__(self,\n                 block,\n                 layers,\n                 block_inplanes,\n                 n_input_channels=1,\n                 conv1_t_size=7,\n                 conv1_t_stride=1,\n                 no_max_pool=False,\n                 shortcut_type='B',\n                 widen_factor=1.0,\n                 n_classes=1):\n        super().__init__()\n\n        block_inplanes = [int(x * widen_factor) for x in block_inplanes]\n\n        self.in_planes = block_inplanes[0]\n        self.no_max_pool = no_max_pool\n\n        self.conv1 = nn.Conv3d(n_input_channels,\n                               self.in_planes,\n                               kernel_size=(conv1_t_size, 7, 7),\n                               stride=(conv1_t_stride, 2, 2),\n                               padding=(conv1_t_size // 2, 3, 3),\n                               bias=False)\n        self.bn1 = nn.BatchNorm3d(self.in_planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool3d(kernel_size=3, stride=2, padding=1)\n        self.layer1 = self._make_layer(block, block_inplanes[0], layers[0],\n                                       shortcut_type)\n        self.layer2 = self._make_layer(block,\n                                       block_inplanes[1],\n                                       layers[1],\n                                       shortcut_type,\n                                       stride=2)\n        self.layer3 = self._make_layer(block,\n                                       block_inplanes[2],\n                                       layers[2],\n                                       shortcut_type,\n                                       stride=2)\n        self.layer4 = self._make_layer(block,\n                                       block_inplanes[3],\n                                       layers[3],\n                                       shortcut_type,\n                                       stride=2)\n\n        self.avgpool = nn.AdaptiveAvgPool3d((1, 1, 1))\n        self.fc = nn.Linear(block_inplanes[3] * block.expansion+1, n_classes)\n        self.un = nn.Linear(block_inplanes[3] * block.expansion+1, n_classes)\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv3d):\n                nn.init.kaiming_normal_(m.weight,\n                                        mode='fan_out',\n                                        nonlinearity='relu')\n            elif isinstance(m, nn.BatchNorm3d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n\n    def _downsample_basic_block(self, x, planes, stride):\n        out = F.avg_pool3d(x, kernel_size=1, stride=stride)\n        zero_pads = torch.zeros(out.size(0), planes - out.size(1), out.size(2),\n                                out.size(3), out.size(4))\n        if isinstance(out.data, torch.cuda.FloatTensor):\n            zero_pads = zero_pads.cuda()\n\n        out = torch.cat([out.data, zero_pads], dim=1)\n\n        return out\n\n    def _make_layer(self, block, planes, blocks, shortcut_type, stride=1):\n        downsample = None\n        if stride != 1 or self.in_planes != planes * block.expansion:\n            if shortcut_type == 'A':\n                downsample = partial(self._downsample_basic_block,\n                                     planes=planes * block.expansion,\n                                     stride=stride)\n            else:\n                downsample = nn.Sequential(\n                    conv1x1x1(self.in_planes, planes * block.expansion, stride),\n                    nn.BatchNorm3d(planes * block.expansion))\n\n        layers = []\n        layers.append(\n            block(in_planes=self.in_planes,\n                  planes=planes,\n                  stride=stride,\n                  downsample=downsample))\n        self.in_planes = planes * block.expansion\n        for i in range(1, blocks):\n            layers.append(block(self.in_planes, planes))\n\n        return nn.Sequential(*layers)\n\n    def forward(self, x, week):\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        if not self.no_max_pool:\n            x = self.maxpool(x)\n\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n\n        x = self.avgpool(x)\n\n        x = x.view(x.size(0), -1)\n        o = torch.sigmoid(self.fc(torch.cat([x, week], dim=1)))\n        u = self.un(torch.cat([x, week], dim=1))\n\n        return o, u\n\ndef generate_model(model_depth, **kwargs):\n    assert model_depth in [10, 18, 34, 50, 101, 152, 200]\n\n    if model_depth == 10:\n        model = ResNet(BasicBlock, [1, 1, 1, 1], get_inplanes(), **kwargs)\n    elif model_depth == 18:\n        model = ResNet(BasicBlock, [2, 2, 2, 2], get_inplanes(), **kwargs)\n    elif model_depth == 34:\n        model = ResNet(BasicBlock, [3, 4, 6, 3], get_inplanes(), **kwargs)\n    elif model_depth == 50:\n        model = ResNet(Bottleneck, [3, 4, 6, 3], get_inplanes(), **kwargs)\n    elif model_depth == 101:\n        model = ResNet(Bottleneck, [3, 4, 23, 3], get_inplanes(), **kwargs)\n    elif model_depth == 152:\n        model = ResNet(Bottleneck, [3, 8, 36, 3], get_inplanes(), **kwargs)\n    elif model_depth == 200:\n        model = ResNet(Bottleneck, [3, 24, 36, 3], get_inplanes(), **kwargs)\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:03:59.799829Z","iopub.execute_input":"2023-01-14T16:03:59.800084Z","iopub.status.idle":"2023-01-14T16:03:59.836699Z","shell.execute_reply.started":"2023-01-14T16:03:59.800060Z","shell.execute_reply":"2023-01-14T16:03:59.835812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net = generate_model(50)","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:03:59.838132Z","iopub.execute_input":"2023-01-14T16:03:59.838686Z","iopub.status.idle":"2023-01-14T16:04:00.567394Z","shell.execute_reply.started":"2023-01-14T16:03:59.838562Z","shell.execute_reply":"2023-01-14T16:04:00.566399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:00.568920Z","iopub.execute_input":"2023-01-14T16:04:00.569295Z","iopub.status.idle":"2023-01-14T16:04:00.606628Z","shell.execute_reply.started":"2023-01-14T16:04:00.569256Z","shell.execute_reply":"2023-01-14T16:04:00.605630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net = net.to(device).float()","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:00.608189Z","iopub.execute_input":"2023-01-14T16:04:00.608855Z","iopub.status.idle":"2023-01-14T16:04:02.221818Z","shell.execute_reply.started":"2023-01-14T16:04:00.608818Z","shell.execute_reply":"2023-01-14T16:04:02.220642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fvc_max = train_csv.FVC.max()\nweek_max = train_csv.Weeks.max()\ntrain_csv.FVC = train_csv.FVC / fvc_max\ntrain_csv.Weeks = train_csv.Weeks / fvc_max","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:02.223545Z","iopub.execute_input":"2023-01-14T16:04:02.224230Z","iopub.status.idle":"2023-01-14T16:04:02.234688Z","shell.execute_reply.started":"2023-01-14T16:04:02.224185Z","shell.execute_reply":"2023-01-14T16:04:02.232009Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def metric(fvc, pred, uncer):\n    sigma_clipped = max(uncer.item(), 70)\n    Delta = min(1000, pred.item() * fvc_max)\n    score = np.sqrt(2) * Delta / sigma_clipped + np.log(np.sqrt(2) * sigma_clipped)\n    return -score","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:02.236639Z","iopub.execute_input":"2023-01-14T16:04:02.237506Z","iopub.status.idle":"2023-01-14T16:04:02.254229Z","shell.execute_reply.started":"2023-01-14T16:04:02.237457Z","shell.execute_reply":"2023-01-14T16:04:02.253044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def UncertLoss(fvc, pred, uncer):\n    first = 1/2 * torch.exp(-uncer) * ((fvc/fvc_max - pred)**2)\n    second = 1/2 * uncer\n    return first + second","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:02.255933Z","iopub.execute_input":"2023-01-14T16:04:02.256723Z","iopub.status.idle":"2023-01-14T16:04:02.268639Z","shell.execute_reply.started":"2023-01-14T16:04:02.256675Z","shell.execute_reply":"2023-01-14T16:04:02.267408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"opt = torch.optim.Adam(net.parameters())","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:02.270433Z","iopub.execute_input":"2023-01-14T16:04:02.271490Z","iopub.status.idle":"2023-01-14T16:04:02.280789Z","shell.execute_reply.started":"2023-01-14T16:04:02.271444Z","shell.execute_reply":"2023-01-14T16:04:02.279620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def windowing(ten, WC=-700, WW=1500):\n    canvas = ten.copy()\n    canvas[ten<WC-WW//2] = WC-WW//2\n    canvas[ten>WC+WW//2] = WC+WW//2\n    canvas = (canvas - canvas.min()) / (canvas.max() - canvas.min())\n    return canvas","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:02.282487Z","iopub.execute_input":"2023-01-14T16:04:02.283450Z","iopub.status.idle":"2023-01-14T16:04:02.292712Z","shell.execute_reply.started":"2023-01-14T16:04:02.283402Z","shell.execute_reply":"2023-01-14T16:04:02.291778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataLoader(torch.utils.data.Dataset):\n    def __init__(self, df, batch_size=8):\n        super().__init__()\n        self.df = df\n        self.bs = batch_size\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        df_shuffled = self.df.sample(frac=1).reset_index(drop=True)\n        ptn = df_shuffled.Patient.tolist()[idx]\n        week = df_shuffled.Weeks.tolist()[idx]\n        fvc = df_shuffled.FVC.tolist()[idx]\n        dcms = []\n        selected = np.random.choice(glob(os.path.join(basepath, 'train', ptn, \"*.dcm\")), self.bs)\n        selected = natsorted(selected)\n        for s in selected:\n            slice_dcm = pydicom.dcmread(s)\n            npy = slice_dcm.pixel_array + float(slice_dcm.RescaleIntercept)\n            dcms.append(windowing(npy))\n        tensor = torch.from_numpy(np.array(dcms)).squeeze().unsqueeze(0).unsqueeze(0)\n        return tensor, torch.from_numpy(np.array([[fvc]])), torch.from_numpy(np.array([[week]]))","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:02.297741Z","iopub.execute_input":"2023-01-14T16:04:02.298820Z","iopub.status.idle":"2023-01-14T16:04:02.310811Z","shell.execute_reply.started":"2023-01-14T16:04:02.298771Z","shell.execute_reply":"2023-01-14T16:04:02.309615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dloader = DataLoader(train_csv)","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:02.315141Z","iopub.execute_input":"2023-01-14T16:04:02.315480Z","iopub.status.idle":"2023-01-14T16:04:02.323729Z","shell.execute_reply.started":"2023-01-14T16:04:02.315449Z","shell.execute_reply":"2023-01-14T16:04:02.322410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"losses = []\nscores = []\n\nfor idx in range(len(train_csv)):\n    try:\n        volume, fvc, week = dloader[idx]\n        volume = volume.to(device).float()\n        fvc = fvc.to(device).float()\n        week = week.to(device).float()\n        pred, uncer = net(volume, week)\n        score = metric(fvc, pred, uncer)\n        loss = UncertLoss(fvc, pred, uncer)\n        opt.zero_grad()\n        loss.backward()\n        opt.step()\n        losses.append(loss.item())\n        scores.append(score.item())\n        print(f\"Step {idx} | Mean Loss {np.mean(losses):.3f} | Mean Score {np.mean(scores):.3f}\")\n        if np.mean(scores)>-6:\n            break\n    except:\n        pass","metadata":{"execution":{"iopub.status.busy":"2023-01-14T16:04:02.326044Z","iopub.execute_input":"2023-01-14T16:04:02.326487Z","iopub.status.idle":"2023-01-14T16:04:16.014125Z","shell.execute_reply.started":"2023-01-14T16:04:02.326446Z","shell.execute_reply":"2023-01-14T16:04:16.013082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}