{"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 torch\nfrom torchvision import datasets\nfrom torchvision import models\nimport torch.nn as nn\nimport os\nfrom skimage import io, transform\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, utils\nimport skimage\nfrom skimage.color import rgb2gray, gray2rgb\nimport cv2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-02-10T07:09:03.153977Z","iopub.execute_input":"2023-02-10T07:09:03.154852Z","iopub.status.idle":"2023-02-10T07:09:06.803846Z","shell.execute_reply.started":"2023-02-10T07:09:03.154746Z","shell.execute_reply":"2023-02-10T07:09:06.802647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def crop_desired_region(img_):\n    coords = cv2.findNonZero(img_) # Find all non-zero points (text)\n    x, y, w, h = cv2.boundingRect(coords) # Find minimum spanning bounding box\n    rect = img_[y:y+h, x:x+w] # Crop the image - note we do this on the original image\n    rect_originalSized = cv2.resize(rect,(img_.shape))\n    return rect_originalSized","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:15:22.966331Z","iopub.execute_input":"2023-02-10T07:15:22.967088Z","iopub.status.idle":"2023-02-10T07:15:22.978994Z","shell.execute_reply.started":"2023-02-10T07:15:22.967043Z","shell.execute_reply":"2023-02-10T07:15:22.97795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def normalize(img: np.ndarray) -> np.ndarray:\n    return (img - img.min()) / (img.max() - img.min())","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:15:23.09983Z","iopub.execute_input":"2023-02-10T07:15:23.101231Z","iopub.status.idle":"2023-02-10T07:15:23.107805Z","shell.execute_reply.started":"2023-02-10T07:15:23.101158Z","shell.execute_reply":"2023-02-10T07:15:23.106289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preproc_tabular(df,root_dir):\n    df['filename'] = root_dir+'/'+df.patient_id.astype(str)+'/'+df.image_id.astype(str)+'.png'\n    train_selected = df.drop(['patient_id','image_id','cancer','filename','laterality','view'],axis=1)\n    train_selected.BIRADS.fillna('7',inplace = True)\n    train_selected.BIRADS = train_selected.BIRADS.astype('category')\n    train_selected.density.fillna('K',inplace = True)\n    train_selected.density = train_selected.density.astype('category')\n    train_input = pd.get_dummies(train_selected)\n    return(train_input)","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:15:23.338791Z","iopub.execute_input":"2023-02-10T07:15:23.339233Z","iopub.status.idle":"2023-02-10T07:15:23.348839Z","shell.execute_reply.started":"2023-02-10T07:15:23.339176Z","shell.execute_reply":"2023-02-10T07:15:23.347891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNADataset(Dataset):\n\n    def __init__(self, csv_file, root_dir, transform=None):\n        self.tabular = pd.read_csv(csv_file)\n        self.root_dir = root_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.tabular)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n        tabular_tmp = self.tabular\n        train_input = preproc_tabular(tabular_tmp,self.root_dir)\n        train_input_tensor = torch.Tensor(train_input.astype('float64').values)\n        img_name = tabular_tmp.filename.tolist()[idx]\n        image = io.imread(img_name)\n        image = crop_desired_region(image)\n        image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)        \n        image = normalize(image)\n        image = torch.from_numpy(image)\n        image = image.permute(2,0,1)\n        tabular_tensor = train_input_tensor[idx]\n        target = torch.Tensor(tabular_tmp['cancer'].tolist())[idx]\n        sample = {'image': image, 'tabular_tensor': tabular_tensor,'answer':target}\n        if self.transform:\n            sample = self.transform(sample)\n        \n        return sample","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:15:23.440509Z","iopub.execute_input":"2023-02-10T07:15:23.440879Z","iopub.status.idle":"2023-02-10T07:15:23.459559Z","shell.execute_reply.started":"2023-02-10T07:15:23.440844Z","shell.execute_reply":"2023-02-10T07:15:23.458583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def stratified_sample_df(df, col, n_samples):\n    n = min(n_samples, df[col].value_counts().min())\n    df_ = df.groupby(col).apply(lambda x: x.sample(n))\n    df_.index = df_.index.droplevel(0)\n    return df_","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:15:23.841696Z","iopub.execute_input":"2023-02-10T07:15:23.842119Z","iopub.status.idle":"2023-02-10T07:15:23.852095Z","shell.execute_reply.started":"2023-02-10T07:15:23.842079Z","shell.execute_reply":"2023-02-10T07:15:23.850926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BasicFC(nn.Module):\n    def __init__(self):\n        super(BasicFC,self).__init__()\n        self.net = nn.Sequential(\n        nn.Linear(16,32),\n        nn.Linear(32,64),\n        nn.Linear(64,128),\n        nn.Linear(128,256),\n        nn.Linear(256,512),\n        nn.Linear(512,1000)\n        )\n    def forward(self,x):\n        return(self.net(x))","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:21:59.712881Z","iopub.execute_input":"2023-02-10T07:21:59.713306Z","iopub.status.idle":"2023-02-10T07:21:59.720193Z","shell.execute_reply.started":"2023-02-10T07:21:59.71327Z","shell.execute_reply":"2023-02-10T07:21:59.719254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PredictCancer(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.img_extractor= models.resnet18(pretrained = True)\n        \n        '''\n        Freeze parameters exclude last fc layer.\n        We're going to train fully connected layer only. \n        '''\n        \n        for params in self.img_extractor.parameters():\n            params.requires_grad = False\n        \n        self.tabular_extractor= BasicFC()\n        self.classifier=nn.Sequential(\n            nn.Linear(in_features= 2000, out_features= 1),\n            nn.SiLU()\n        )\n    def forward(self, image, tabular):\n        image_feature= self.img_extractor(image.float())\n        tabular_feature= self.tabular_extractor(tabular)\n        output= torch.cat([image_feature, tabular_feature], dim=-1)\n        output= self.classifier(output)\n        return output","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:22:00.735034Z","iopub.execute_input":"2023-02-10T07:22:00.735403Z","iopub.status.idle":"2023-02-10T07:22:00.743031Z","shell.execute_reply.started":"2023-02-10T07:22:00.735371Z","shell.execute_reply":"2023-02-10T07:22:00.742077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32\ncsv ='/kaggle/input/rsna-breast-cancer-detection/train.csv'\nstratified_sample_df(pd.read_csv(csv),'cancer',10000).reset_index(drop = True).to_csv('./modified.csv',index = False)\ncsv ='./modified.csv'\n\nrootpath = '/kaggle/input/rsna-mammography-images-as-pngs/images_as_pngs_512/train_images_processed_512/'\nrsna_dataset = RSNADataset(csv_file=csv,root_dir=rootpath)\nrsna_dataloader = DataLoader(rsna_dataset, batch_size=batch_size, shuffle=True, num_workers=0)\nmodel = PredictCancer()\n","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:22:00.870244Z","iopub.execute_input":"2023-02-10T07:22:00.87059Z","iopub.status.idle":"2023-02-10T07:22:01.154896Z","shell.execute_reply.started":"2023-02-10T07:22:00.870561Z","shell.execute_reply":"2023-02-10T07:22:01.153898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nfrom torch import optim\ndevice = 'cuda'\nmodel = model.cuda()\nlr = 1e-05\nnum_epochs = 1\noptimizer = optim.Adam(model.parameters(), lr=lr)\nloss_function = nn.CrossEntropyLoss().to(device)\nloss_function = nn.BCEWithLogitsLoss()\nparams = {\n    'num_epochs':num_epochs,\n    'optimizer':optimizer,\n    'loss_function':loss_function,\n    'device':device\n}","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:22:01.397815Z","iopub.execute_input":"2023-02-10T07:22:01.398162Z","iopub.status.idle":"2023-02-10T07:22:01.422899Z","shell.execute_reply.started":"2023-02-10T07:22:01.398133Z","shell.execute_reply":"2023-02-10T07:22:01.421952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, params):\n    loss_function=params[\"loss_function\"]\n    device=params[\"device\"]\n    model.train()\n    for epoch in range(0, num_epochs):\n        for dat in rsna_dataloader:\n            image_batch = dat['image'].cuda()\n            tabular_batch = dat['tabular_tensor'].cuda()\n            y_batch = dat['answer'].cuda()\n            optimizer.zero_grad() \n            outputs = model(image_batch,tabular_batch)\n            outputs = outputs.squeeze()\n            train_separate_loss = loss_function(outputs,y_batch).requires_grad_(True)\n            train_separate_loss.backward()\n            optimizer.step()\n            print('Epoch: %d/%d, Train loss: %.6f' %(epoch+1, num_epochs, train_separate_loss.item()))\n\ntrain_model(model,params)","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:22:02.448168Z","iopub.execute_input":"2023-02-10T07:22:02.448927Z","iopub.status.idle":"2023-02-10T07:23:39.533857Z","shell.execute_reply.started":"2023-02-10T07:22:02.448891Z","shell.execute_reply":"2023-02-10T07:23:39.53177Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model","metadata":{"execution":{"iopub.status.busy":"2023-02-10T07:23:58.270566Z","iopub.execute_input":"2023-02-10T07:23:58.270945Z","iopub.status.idle":"2023-02-10T07:23:58.280785Z","shell.execute_reply.started":"2023-02-10T07:23:58.270912Z","shell.execute_reply":"2023-02-10T07:23:58.279547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}