{"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":"from matplotlib import pyplot as plt\nimport pandas as pd\nimport os\nimport torchvision.transforms as tr\nimport numpy as np\nimport torch\nfrom torch import nn\nfrom torch import optim\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport pydicom\nimport glob\nimport collections\nfrom datetime import datetime\nfrom skimage import measure\nfrom skimage.measure import block_reduce\nfrom matplotlib import pyplot as plt\nfrom mpl_toolkits.mplot3d.art3d import Poly3DCollection\nfrom skimage.transform import resize\nimport torchvision.models as models\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.autograd import Variable\nimport seaborn as sns\nfrom collections import defaultdict","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-19T05:34:25.494603Z","iopub.execute_input":"2023-01-19T05:34:25.495012Z","iopub.status.idle":"2023-01-19T05:34:30.868755Z","shell.execute_reply.started":"2023-01-19T05:34:25.494923Z","shell.execute_reply":"2023-01-19T05:34:30.8679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install torch==1.11.0+cu113 torchvision==0.12.0+cu113 torchaudio==0.11.0 --extra-index-url https://download.pytorch.org/whl/cu113","metadata":{"execution":{"iopub.status.busy":"2023-01-19T04:38:42.343092Z","iopub.execute_input":"2023-01-19T04:38:42.343455Z","iopub.status.idle":"2023-01-19T04:38:46.4189Z","shell.execute_reply.started":"2023-01-19T04:38:42.343422Z","shell.execute_reply":"2023-01-19T04:38:46.417772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:34:35.329707Z","iopub.execute_input":"2023-01-19T05:34:35.330076Z","iopub.status.idle":"2023-01-19T05:34:35.334277Z","shell.execute_reply.started":"2023-01-19T05:34:35.330042Z","shell.execute_reply":"2023-01-19T05:34:35.333332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('../input/vinbigdata-chest-xray-abnormalities-detection/train.csv')\ndf.head(10)","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:34:36.46194Z","iopub.execute_input":"2023-01-19T05:34:36.462267Z","iopub.status.idle":"2023-01-19T05:34:36.658579Z","shell.execute_reply.started":"2023-01-19T05:34:36.462232Z","shell.execute_reply":"2023-01-19T05:34:36.657872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_file = os.listdir('../input/vinbigdata-chest-xray-abnormalities-detection/train')\ninput_files = []\nfor ip in input_file:\n    input_files.append(ip.split('.')[0])","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:34:37.389924Z","iopub.execute_input":"2023-01-19T05:34:37.390235Z","iopub.status.idle":"2023-01-19T05:34:37.733579Z","shell.execute_reply.started":"2023-01-19T05:34:37.390204Z","shell.execute_reply":"2023-01-19T05:34:37.73221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(input_files)","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:34:38.14523Z","iopub.execute_input":"2023-01-19T05:34:38.145571Z","iopub.status.idle":"2023-01-19T05:34:38.151137Z","shell.execute_reply.started":"2023-01-19T05:34:38.145539Z","shell.execute_reply":"2023-01-19T05:34:38.150274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[df['image_id']==input_files[3]].sort_values(by=['class_id'])","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:34:39.103139Z","iopub.execute_input":"2023-01-19T05:34:39.103467Z","iopub.status.idle":"2023-01-19T05:34:39.128048Z","shell.execute_reply.started":"2023-01-19T05:34:39.103436Z","shell.execute_reply":"2023-01-19T05:34:39.127213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class traindataset(torch.utils.data.Dataset):\n\n    def __init__(self, df, file_list, transform):\n        super().__init__()\n        \n        self.file_list = file_list\n        self.df = df\n    \n    def __len__(self) :\n        return len(self.file_list)\n\n    def __getitem__(self, idx):\n        \n        img_id = self.file_list[idx]\n        df = self.df\n        s = 224\n        N = 15\n        d_f = df[df['image_id']==img_id]\n        dff = d_f.sort_values(by=['class_id'])\n        class_id = dff['class_id'].values.tolist()\n        img_pxl = pydicom.read_file('../input/vinbigdata-chest-xray-abnormalities-detection/train/'+img_id+'.dicom').pixel_array\n        img_res = resize(img_pxl,(s,s),anti_aliasing=True)\n        img_np = img_res.astype(np.float32())\n        img_tr = torch.from_numpy(img_np)\n        x_ = s/img_pxl.shape[1]\n        y_ = s/img_pxl.shape[0]\n        xmin = [x*x_ for x in dff['x_min'].values.tolist()]\n        ymin = [y*y_ for y in dff['y_min'].values.tolist()]\n        xmax = [x1*x_ for x1 in dff['x_max'].values.tolist()]\n        ymax = [y1*y_ for y1 in dff['y_max'].values.tolist()]\n        #bbox = []\n        #for z in range(len(xmin)):\n        #    bbox.append([xmin[z],ymin[z],xmax[z],ymax[z]])\n        mask = np.zeros((N,s,s))\n        for k,m in enumerate(class_id):\n            if m != 14:\n                x1,x2,y1,y2 = int(xmin[k]),int(xmax[k]),int(ymin[k]),int(ymax[k])\n                mask[m,y1:y2,x1:x2] = 1\n        mask_numpy = mask.astype(np.float32())\n        mask_tensor = torch.from_numpy(mask_numpy)\n        #mask_tensor = np.transpose(mask_tensor, (2,0,1))\n        return img_tr,mask_tensor","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:34:39.55004Z","iopub.execute_input":"2023-01-19T05:34:39.550365Z","iopub.status.idle":"2023-01-19T05:34:39.562252Z","shell.execute_reply.started":"2023-01-19T05:34:39.550332Z","shell.execute_reply":"2023-01-19T05:34:39.561096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"traindata = traindataset(file_list = input_files,df =df,transform = None)\ndata_loader = torch.utils.data.DataLoader(traindata, batch_size=1, shuffle=True, num_workers=1)","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:34:40.164764Z","iopub.execute_input":"2023-01-19T05:34:40.165139Z","iopub.status.idle":"2023-01-19T05:34:40.169588Z","shell.execute_reply.started":"2023-01-19T05:34:40.165107Z","shell.execute_reply":"2023-01-19T05:34:40.168619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib.patches import Rectangle\nfor k,kk in enumerate(traindata):\n    plt.imshow(kk[0])\n    \n        #plt.gca().add_patch(Rectangle((org[0], org[1]), (org[2]-org[0]), (org[3]-org[1]),linewidth=1,edgecolor='b',facecolor='none'))\n \n    for j in kk[1]:\n       \n        plt.gca().add_patch(Rectangle((org[0], org[1]), (org[2]-org[0]), (org[3]-org[1]),linewidth=1,edgecolor='b',facecolor='none'))\n \n        plt.imshow(j)\n        plt.show()\n    if k == 4:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:34:40.875411Z","iopub.execute_input":"2023-01-19T05:34:40.875828Z","iopub.status.idle":"2023-01-19T05:35:04.901345Z","shell.execute_reply.started":"2023-01-19T05:34:40.875772Z","shell.execute_reply":"2023-01-19T05:35:04.900645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:35:43.502575Z","iopub.execute_input":"2023-01-19T05:35:43.502918Z","iopub.status.idle":"2023-01-19T05:35:43.507064Z","shell.execute_reply.started":"2023-01-19T05:35:43.502885Z","shell.execute_reply":"2023-01-19T05:35:43.505909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class batchnorm_relu(nn.Module):\n    def __init__(self, in_c):\n        super().__init__()\n\n        self.bn = nn.BatchNorm2d(in_c)\n        self.relu = nn.ReLU()\n\n    def forward(self, inputs):\n        x = self.bn(inputs)\n        x = self.relu(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:35:44.892602Z","iopub.execute_input":"2023-01-19T05:35:44.892948Z","iopub.status.idle":"2023-01-19T05:35:44.898161Z","shell.execute_reply.started":"2023-01-19T05:35:44.892901Z","shell.execute_reply":"2023-01-19T05:35:44.897289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class residual_block(nn.Module):\n    def __init__(self, in_c, out_c, stride=1):\n        super().__init__()\n\n        \"\"\" Convolutional layer \"\"\"\n        self.b1 = batchnorm_relu(in_c)\n        self.c1 = nn.Conv2d(in_c, out_c, kernel_size=3, padding=1, stride=stride)\n        self.b2 = batchnorm_relu(out_c)\n        self.c2 = nn.Conv2d(out_c, out_c, kernel_size=3, padding=1, stride=1)\n\n        \"\"\" Shortcut Connection (Identity Mapping) \"\"\"\n        self.s = nn.Conv2d(in_c, out_c, kernel_size=1, padding=0, stride=stride)\n\n    def forward(self, inputs):\n        x = self.b1(inputs)\n        x = self.c1(x)\n        x = self.b2(x)\n        x = self.c2(x)\n        s = self.s(inputs)\n\n        skip = x + s\n        return skip","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:35:45.932773Z","iopub.execute_input":"2023-01-19T05:35:45.933128Z","iopub.status.idle":"2023-01-19T05:35:45.93991Z","shell.execute_reply.started":"2023-01-19T05:35:45.933095Z","shell.execute_reply":"2023-01-19T05:35:45.939053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class decoder_block(nn.Module):\n    def __init__(self, in_c, out_c):\n        super().__init__()\n\n        self.upsample = nn.Upsample(scale_factor=2, mode=\"bilinear\", align_corners=True)\n        self.r = residual_block(in_c+out_c, out_c)\n\n    def forward(self, inputs, skip):\n        x = self.upsample(inputs)\n        x = torch.cat([x, skip], axis=1)\n        x = self.r(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:35:46.569912Z","iopub.execute_input":"2023-01-19T05:35:46.570238Z","iopub.status.idle":"2023-01-19T05:35:46.57771Z","shell.execute_reply.started":"2023-01-19T05:35:46.570207Z","shell.execute_reply":"2023-01-19T05:35:46.576954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class build_resunet(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        \"\"\" Encoder 1 \"\"\"\n        self.c11 = nn.Conv2d(1, 64, kernel_size=3, padding=1)\n        self.br1 = batchnorm_relu(64)\n        self.c12 = nn.Conv2d(64, 64, kernel_size=3, padding=1)\n        self.c13 = nn.Conv2d(1, 64, kernel_size=1, padding=0)\n\n        \"\"\" Encoder 2 and 3 \"\"\"\n        self.r2 = residual_block(64, 128, stride=2)\n        self.r3 = residual_block(128, 256, stride=2)\n\n        \"\"\" Bridge \"\"\"\n        self.r4 = residual_block(256, 512, stride=2)\n\n        \"\"\" Decoder \"\"\"\n        self.d1 = decoder_block(512, 256)\n        self.d2 = decoder_block(256, 128)\n        self.d3 = decoder_block(128, 64)\n\n        \"\"\" Output \"\"\"\n        self.output = nn.Conv2d(64, 1, kernel_size=1, padding=0)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, inputs):\n        \"\"\" Encoder 1 \"\"\"\n        x = self.c11(inputs)\n        x = self.br1(x)\n        x = self.c12(x)\n        s = self.c13(inputs)\n        skip1 = x + s\n\n        \"\"\" Encoder 2 and 3 \"\"\"\n        skip2 = self.r2(skip1)\n        skip3 = self.r3(skip2)\n\n        \"\"\" Bridge \"\"\"\n        b = self.r4(skip3)\n\n        \"\"\" Decoder \"\"\"\n        d1 = self.d1(b, skip3)\n        d2 = self.d2(d1, skip2)\n        d3 = self.d3(d2, skip1)\n\n        \"\"\" output \"\"\"\n        output = self.output(d3)\n        output = self.sigmoid(output)\n\n        return output","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:37:48.909921Z","iopub.execute_input":"2023-01-19T05:37:48.910268Z","iopub.status.idle":"2023-01-19T05:37:48.920824Z","shell.execute_reply.started":"2023-01-19T05:37:48.910235Z","shell.execute_reply":"2023-01-19T05:37:48.919808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_unet = build_resunet()\nmodel_unet = model_unet.cuda()\n\ncalc = nn.MSELoss()\noptimizer = optim.Adamax(model_unet.parameters(), lr=0.0003)\n","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:37:49.926454Z","iopub.execute_input":"2023-01-19T05:37:49.926769Z","iopub.status.idle":"2023-01-19T05:37:50.014872Z","shell.execute_reply.started":"2023-01-19T05:37:49.926737Z","shell.execute_reply":"2023-01-19T05:37:50.014064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_loss(pred, target, smooth = 1.):\n    pred = pred.contiguous()\n    target = target.contiguous()    \n\n    intersection = (pred * target).sum(dim=2).sum(dim=2)\n    \n    loss = (1 - ((2. * intersection + smooth) / (pred.sum(dim=2).sum(dim=2) + target.sum(dim=2).sum(dim=2) + smooth)))\n    \n    return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:37:52.218682Z","iopub.execute_input":"2023-01-19T05:37:52.219049Z","iopub.status.idle":"2023-01-19T05:37:52.226074Z","shell.execute_reply.started":"2023-01-19T05:37:52.219016Z","shell.execute_reply":"2023-01-19T05:37:52.224949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calc_loss(pred, target, metrics, bce_weight=0.5):\n    bce = F.binary_cross_entropy_with_logits(pred, target)\n        \n    pred = F.sigmoid(pred)\n    dice = dice_loss(pred, target)\n    \n    loss = bce * bce_weight + dice * (1 - bce_weight)\n    \n    metrics['bce'] += bce.data.cpu().numpy() * target.size(0)\n    metrics['dice'] += dice.data.cpu().numpy() * target.size(0)\n    metrics['loss'] += loss.data.cpu().numpy() * target.size(0)\n    \n    return loss","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:37:52.74259Z","iopub.execute_input":"2023-01-19T05:37:52.742932Z","iopub.status.idle":"2023-01-19T05:37:52.749141Z","shell.execute_reply.started":"2023-01-19T05:37:52.7429Z","shell.execute_reply":"2023-01-19T05:37:52.74808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics = defaultdict(float)","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:37:53.224296Z","iopub.execute_input":"2023-01-19T05:37:53.224625Z","iopub.status.idle":"2023-01-19T05:37:53.229084Z","shell.execute_reply.started":"2023-01-19T05:37:53.224594Z","shell.execute_reply":"2023-01-19T05:37:53.22786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = 1\nsteps = 0\nprint_every = 750\ntrain_losses, train_accuracy = [], []\n#model.load_state_dict(torch.load('./final_model.pth'))\nfor epoch in range(epochs):\n    model_unet.train()\n    size = 0\n    running_loss = 0\n    acc = 0\n    for a,(image_train, y_train) in enumerate(data_loader):\n        steps += 1\n        image_train, y_train = image_train.unsqueeze(0).cuda(), y_train.cuda()\n        image_train = Variable(image_train,requires_grad=True)\n        #image_train=  image_train.detach().cpu().numpy()\n        #image_train = np.dstack([image_train]*3)\n        #image_train = image_train.reshape(-1,3,224,224)\n        #image_train = torch.from_numpy(image_train).cuda()\n        \n        optimizer.zero_grad()\n        y_predtrain = model_unet.forward(image_train)\n        #y_train=y_train.type(torch.LongTensor)\n        #loss = calc_loss(y_predtrain, y_train,metrics)\n        loss = calc(y_predtrain, y_train)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item()\n        ps = torch.exp(y_predtrain)\n        top_p = torch.max(y_predtrain, 1)\n        top_class = torch.argmax(y_predtrain,dim = 1)\n        equals = top_class == y_train\n        print(torch.mean(equals.type(torch.FloatTensor)).item())\n        acc += torch.mean(equals.type(torch.FloatTensor)).item()\n        size += image_train.shape[0]\n        model_unet.eval()\n        print(f\"Epoch {epoch+1}/{epochs}.. \"\n              f\"Train loss: {running_loss/print_every:.3f}.. \"\n              f\"Train accuracy: {acc/len(data_loader):.3f}\")\n    #torch.save(model.state_dict(),'./'+str(epoch)+'model.pth')\n    #print('model saved')\n    torch.save(model_unet.state_dict(),'./'+str(epoch)+'unet_model.pth')\n    train_losses.append(float(running_loss)/float(size))\n    train_accuracy.append(float(acc)/float(size))\n    print('train_losses',epoch,train_losses)\n    print('train_accuracy',epoch,train_accuracy)\ntorch.save(model_unet.state_dict(),'./final_unet_model.pth')\nprint('model saved')\nprint('train_losses',epoch,train_losses)\nprint('train_accuracy',epoch,train_accuracy)","metadata":{"execution":{"iopub.status.busy":"2023-01-19T05:37:53.828267Z","iopub.execute_input":"2023-01-19T05:37:53.82859Z","iopub.status.idle":"2023-01-19T07:37:43.092247Z","shell.execute_reply.started":"2023-01-19T05:37:53.82856Z","shell.execute_reply":"2023-01-19T07:37:43.088434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Great!! Got an accuracy of almost 99.7 percent!! I think the model is overfitting.","metadata":{}},{"cell_type":"code","source":"img_pxl = pydicom.read_file('../input/vinbigdata-chest-xray-abnormalities-detection/train/000d68e42b71d3eac10ccc077aba07c1.dicom').pixel_array\nimg_res = resize(img_pxl,(224,224),anti_aliasing=True)\nimg_np = img_res.astype(np.float32())\nimg_tr = torch.from_numpy(img_np).unsqueeze(0).unsqueeze(0).cuda()\nimg_tr = Variable(img_tr,requires_grad=True)\noutput = model_unet(img_tr)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict = output.cpu().detach().squeeze()\nprint(predict.shape)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in predict:\n    plt.imshow(i)\n    plt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}