{"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":"markdown","source":"# 1.Import module","metadata":{}},{"cell_type":"code","source":"# # !pip install efficientnet_pytorch\n!pip install torchsummary\n# try:\n#     import pylibjpeg\n# except:\n#     !rm -rf /root/.cache/torch/hub/checkpoints/\n#     !mkdir -p /root/.cache/torch/hub/checkpoints/\n#     !pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}\n#     !pip install /kaggle/input/rsna-2022-whl/{torch-1.12.1-cp37-cp37m-manylinux1_x86_64.whl,torchvision-0.13.1-cp37-cp37m-manylinux1_x86_64.whl}","metadata":{"execution":{"iopub.status.busy":"2023-01-15T13:49:59.87015Z","iopub.execute_input":"2023-01-15T13:49:59.87109Z","iopub.status.idle":"2023-01-15T13:50:11.354873Z","shell.execute_reply.started":"2023-01-15T13:49:59.8709Z","shell.execute_reply":"2023-01-15T13:50:11.353654Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor\nimport re\nimport matplotlib.pyplot as plt\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\nimport torchvision\nimport pydicom\nfrom glob import glob\nfrom tqdm.notebook import tqdm\nimport cv2\n#visualize\nimport scipy\nimport seaborn as sns\n# PyTorch\nimport torch\nfrom skimage import filters\nimport pytorch_lightning as pl\nimport torchvision\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch import FloatTensor, LongTensor\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\n# Data Augmentation for Image Preprocessing\nfrom albumentations import (ToFloat, Normalize, VerticalFlip, HorizontalFlip, Compose, Resize,\n                            RandomBrightnessContrast, HueSaturationValue, Blur, GaussNoise,\n                            Rotate, RandomResizedCrop, Cutout, ShiftScaleRotate, ToGray)\nfrom albumentations.pytorch import ToTensorV2\n\n# from efficientnet_pytorch import EfficientNet\n# from torchvision.models import resnet34, resnet50\n\nfrom torchsummary import summary\n\n# SKlearn\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold\nfrom sklearn.metrics import accuracy_score, roc_auc_score, confusion_matrix\n\n#Global val\nclass _color:\n    S = '\\033[1m' + '\\033[92m'\n    E = '\\033[0m'\n    \nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(_color.S+'Device available now:'+_color.E, DEVICE)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-01-15T13:50:11.358133Z","iopub.execute_input":"2023-01-15T13:50:11.359281Z","iopub.status.idle":"2023-01-15T13:50:17.419242Z","shell.execute_reply.started":"2023-01-15T13:50:11.359238Z","shell.execute_reply":"2023-01-15T13:50:17.418223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### WandB login","metadata":{}},{"cell_type":"code","source":"# 🐝 Secrets\nimport wandb\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"wandb\")\n! wandb login $secret_value_0","metadata":{"execution":{"iopub.status.busy":"2023-01-15T13:50:17.422555Z","iopub.execute_input":"2023-01-15T13:50:17.42285Z","iopub.status.idle":"2023-01-15T13:50:19.951211Z","shell.execute_reply.started":"2023-01-15T13:50:17.422807Z","shell.execute_reply":"2023-01-15T13:50:19.949985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ntest = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/test.csv')\ntrain.head(3)","metadata":{"execution":{"iopub.status.busy":"2023-01-15T13:50:19.955847Z","iopub.execute_input":"2023-01-15T13:50:19.956156Z","iopub.status.idle":"2023-01-15T13:50:20.087632Z","shell.execute_reply.started":"2023-01-15T13:50:19.956127Z","shell.execute_reply":"2023-01-15T13:50:20.086636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test.head()","metadata":{"execution":{"iopub.status.busy":"2023-01-15T13:50:20.089034Z","iopub.execute_input":"2023-01-15T13:50:20.089691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.iloc[0].cancer==0","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**First We will train model trainby custom RSNA dataset we has croped**<br>\nAnd we has found a little bug at index 53065 so we will drop it out from our Custom Dataset","metadata":{}},{"cell_type":"markdown","source":"**About our custom dataset : it been croped from 256px PNG dataset**","metadata":{}},{"cell_type":"code","source":"print(_color.S+'Image Bug frame :\\n'+_color.E+str(train.iloc[53065]))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train=train.drop(53065) #bug in index 53065\nbase_path='/kaggle/input/rsna-256px-croped/RSNA_Crop'\nall_paths = []\nfor k in tqdm(range(len(train))):\n    if train.iloc[k].cancer==0:\n        row = train.iloc[k, :]\n        all_paths.append(base_path + '/nocancer/'+str(row.patient_id) + \"_\" + str(row.image_id) + \".png\")\n    else:\n        row = train.iloc[k, :]\n        all_paths.append(base_path + '/cancer/'+str(row.patient_id) + \"_\" + str(row.image_id) + \".png\")\ntrain[\"path\"] = all_paths","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Save new metadata","metadata":{"execution":{"iopub.status.busy":"2023-01-15T04:47:17.332451Z","iopub.execute_input":"2023-01-15T04:47:17.333143Z","iopub.status.idle":"2023-01-15T04:47:17.339569Z","shell.execute_reply.started":"2023-01-15T04:47:17.333104Z","shell.execute_reply":"2023-01-15T04:47:17.338093Z"}}},{"cell_type":"code","source":"train.to_csv('RSNA_metadata.csv',index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv('/kaggle/working/RSNA_metadata.csv',index_col=0)\ndf_train.iloc[[53065]]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Read_dicom and crop","metadata":{}},{"cell_type":"markdown","source":"**Q**: Why we have to read and crop those image<br>\n**A**: Because we only croped those image for train data so now we will processcing a litte bit for test set competition","metadata":{"execution":{"iopub.status.busy":"2023-01-15T05:00:00.609822Z","iopub.execute_input":"2023-01-15T05:00:00.610254Z","iopub.status.idle":"2023-01-15T05:00:00.621163Z","shell.execute_reply.started":"2023-01-15T05:00:00.610212Z","shell.execute_reply":"2023-01-15T05:00:00.619747Z"}}},{"cell_type":"markdown","source":"**Yeah we found does crop by conected_components on [this notebook](https://www.kaggle.com/code/alimbekovkz/eda-image-crop-albumentations-augs) so we reference and improve it a litte bit**\n","metadata":{}},{"cell_type":"code","source":"def get_masks_and_sizes_of_connected_components(img_mask):\n    \"\"\"\n    Finds the connected components from the mask of the image\n    \"\"\"\n    mask, num_labels = scipy.ndimage.label(img_mask)\n\n    mask_pixels_dict = {}\n    for i in range(num_labels+1):\n        this_mask = (mask == i)\n        if img_mask[this_mask][0] != 0:\n            # Exclude the 0-valued mask\n            mask_pixels_dict[i] = np.sum(this_mask)\n        \n    return mask, mask_pixels_dict\n\n\ndef get_mask_of_largest_connected_component(img_mask):\n    \"\"\"\n    Finds the largest connected component from the mask of the image\n    \"\"\"\n    mask, mask_pixels_dict = get_masks_and_sizes_of_connected_components(img_mask)\n    largest_mask_index = pd.Series(mask_pixels_dict).idxmax()\n    largest_mask = mask == largest_mask_index\n    return largest_mask\n\ndef image_procescing(img):\n    \"\"\"\n    Crop image by find coordinates of the largest connected componen\n    \"\"\"\n    #check_img_convert_gray\n    if len(img.shape)==3:\n        img = rgb2gray(img)\n    #convert to bin\n    \n    threshold = filters.threshold_isodata(img)\n    bin_img = (img > threshold)*1\n    kernel = np.ones((5, 5), np.uint8)\n    bin_img = bin_img.astype('uint8')\n    bin_img = cv2.erode(bin_img, kernel, iterations=-2)\n    \n    #most mask\n    img_mask = get_mask_of_largest_connected_component(bin_img)\n    #crop_image\n    \n    farest_pixel = np.max(list(zip(*np.where(img_mask == 1))), axis=0)\n    nearest_pixel = np.min(list(zip(*np.where(img_mask == 1))), axis=0)\n    croped =  img[nearest_pixel[0]:farest_pixel[0], nearest_pixel[1]:farest_pixel[1]]\n    return croped","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Read_dicom_to512 \nThis read_dicom Function we reference on top 3 voted notebook of this competetion rightnow [here is this notebook](https://www.kaggle.com/code/theoviel/dicom-resized-png-jpg)","metadata":{}},{"cell_type":"code","source":"def read_dicom_512(f, size=512):\n    \"\"\"\n    Read dicom path\n    \"\"\"\n    dicom = pydicom.dcmread(f)\n    img = dicom.pixel_array\n    img = (img - img.min()) / (img.max() - img.min())\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        img = 1 - img\n\n    img = cv2.resize(img, (size, size))\n    \n    return img","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset,Dataloader","metadata":{}},{"cell_type":"markdown","source":"We Decide to use pytorch to train this notebook and here is Dataset of our model","metadata":{}},{"cell_type":"markdown","source":"Our input image will be a tensor 1x227x227 so we'll Transform it ","metadata":{}},{"cell_type":"code","source":"class RSNA_dataset(Dataset):\n    def __init__(self,csv_file,root_dir,transform=None,is_train = True):\n        \"\"\"\n        Read metadata, root for image, and check if it is train or not\n        \"\"\"\n        self.dataframe = pd.read_csv(csv_file)\n        self.root_dir = root_dir\n        self.is_train = is_train\n        \n        if is_train:\n            self.transform = transform\n        else:\n            self.transform = Compose([Resize(height=227,width=227,always_apply=True),\n                                      ToTensorV2()])\n    \n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self,index):\n        \"\"\"\n        Read and transform if it is trainset and preprocescing if it is test set\n        \"\"\"\n        if self.is_train:\n        #data_frame\n            if self.dataframe.iloc[index]['cancer']==0:\n                image_path = self.root_dir + '/nocancer/'+str(self.dataframe.iloc[index].patient_id) + \"_\" + str(self.dataframe.iloc[index].image_id) + \".png\"\n            else:\n                image_path = self.root_dir + '/cancer/'+str(self.dataframe.iloc[index].patient_id) + \"_\" + str(self.dataframe.iloc[index].image_id) + \".png\"\n            #read image\n            image = cv2.imread(image_path,0)\n            #transform image\n            if self.transform != None:\n                image_trans = self.transform(image=image)['image']\n                label = self.dataframe.iloc[index]['cancer'] \n            else:\n                image_trans = image\n                label = self.dataframe.iloc[index]['cancer'] \n            return image_trans,label\n        \n        else: \n            image_path = self.root_dir+'/'+str(self.dataframe.iloc[index].patient_id) + \"/\" + str(self.dataframe.iloc[index].image_id) + \".dcm\"\n            image = read_dicom_512(image_path)\n            img = image_procescing(image)\n            image_trans = self.transform(image=img)['image']\n            prediction_id = self.dataframe.iloc[index]['prediction_id']\n            return image_trans,prediction_id","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we will define path of our dataset, and write a Compose transform for train set","metadata":{}},{"cell_type":"code","source":"csv_path = '/kaggle/input/breastcancer/RSNA_metadata.csv'\nimage_path = '/kaggle/input/rsna-256px-croped/RSNA_Crop'\ntest_csv_path='/kaggle/input/rsna-breast-cancer-detection/test.csv'\ntest_image_path='/kaggle/input/rsna-breast-cancer-detection/test_images'\n\ntransform = Compose([Resize(height=227,width=227,always_apply=True),\n                    Normalize(mean=0.449,std=0.226),\n                    HorizontalFlip(),\n                    VerticalFlip(),\n                    Rotate(),\n#                     RandomBrightnessContrast(p=0.15),\n                    ToTensorV2()])\n\ndef data_to_device(img,label=None):\n    if label !=None:\n        return img.to(DEVICE), label.to(DEVICE)\n    else:\n        return img.to(DEVICE)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**NOTE** : remember to change those tensor to CUDA :3","metadata":{}},{"cell_type":"code","source":"train_dataset=RSNA_dataset(csv_path,image_path,transform)\ntrain_dataloader = DataLoader(train_dataset, batch_size=16, shuffle=True)\ntest_dataset=RSNA_dataset(test_csv_path,test_image_path,is_train = False)\ntest_dataloader = DataLoader(test_dataset, batch_size=4, shuffle=False)\nfor k,(img,la) in enumerate(train_dataloader):\n    img,la = data_to_device(img,la)\n    print(_color.S + f\"Batch: {k}\" + _color.E, \"\\n\" +\n          _color.S + \"Image:\" + _color.E, img.shape, \"\\n\" +\n          _color.S + \"Label:\" + _color.E, la, \"\\n\" +\n          \"=\"*50)\n    if k==4:break\nfor k,(img,pred_id) in enumerate(test_dataloader):\n    print(_color.S + f\"Batch: {k}\" + _color.E, \"\\n\" +\n          _color.S + \"Image:\" + _color.E, img.shape, \"\\n\" +\n          _color.S +str(pred_id)+'\\n'+_color.E+\n          \"=\"*50)\n    if k==4:break","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_t,la=train_dataset.__getitem__(5)\nplt.imshow(img_t.permute(1,2,0),cmap=plt.cm.gray);\nprint(_color.S+'Label of this image is:'+_color.E,la)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_test,pred_id=test_dataset.__getitem__(0)\nplt.imshow(img_test.permute(1,2,0),cmap=plt.cm.gray);\nprint(_color.S+'Sample of image Test is:'+_color.E)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" # Model CNN ","metadata":{}},{"cell_type":"markdown","source":"We reference CNN model structure from [this paper](https://drive.google.com/file/d/1DHvb2ssSzN5zcXPlPC1zMh_2F1Pr7Slb/view)<br>\nThose look like this:\n![image](https://scontent.fsgn4-1.fna.fbcdn.net/v/t1.15752-9/317923657_2181252688743181_864461044515013559_n.png?_nc_cat=101&ccb=1-7&_nc_sid=ae9488&_nc_ohc=FiMCnqGKQgsAX89rEOe&_nc_ht=scontent.fsgn4-1.fna&oh=03_AdQTD77LNTfWxk_oQwiEmq0rq1njUA5E0myPfIkNalVxGA&oe=63EB0011)","metadata":{}},{"cell_type":"code","source":"class CNN(nn.Module):\n    def __init__(self):\n        super(CNN,self).__init__()\n        def layer(input_channel: int, \n                  output_channel: int, \n                  kernel_size_conv: int=3, \n                  padding: str='same', \n                  stride_conv: int=1, \n                  kernel_size_maxpool: int=2, \n                  stride_maxpool: int=2,\n                  batchnorm=True,\n                  maxpool=True,):\n            layers = [nn.Conv2d(input_channel, output_channel, kernel_size_conv,stride_conv, padding)]\n            if batchnorm:\n                layers.append(nn.BatchNorm2d(output_channel))\n                \n            layers.append(nn.ReLU())\n            \n            if maxpool:\n                layers.append(nn.MaxPool2d(kernel_size_maxpool, stride_maxpool))\n                \n            return layers\n        \n        self.conv_seq = nn.Sequential(\n                        *layer(input_channel=1, \n                               output_channel=8 ,\n                               kernel_size_conv=3, \n                               stride_conv=1,\n                               padding='same',\n                               kernel_size_maxpool=2,\n                               stride_maxpool=2),\n            \n                        *layer(input_channel=8, \n                               output_channel=16 ,\n                               kernel_size_conv=3, \n                               stride_conv=1,\n                               padding='same',\n                               kernel_size_maxpool=2,\n                               stride_maxpool=2),\n            \n                        *layer(input_channel=16, \n                               output_channel=32,\n                               kernel_size_conv=3, \n                               stride_conv=1,\n                               padding='same',\n                               maxpool=False),\n                    )\n        self.fc_seq = nn.Sequential(nn.Linear(100352,512),\n                                   nn.ReLU(),\n                                   nn.Linear(512,128),\n                                   nn.ReLU(),\n                                   nn.Linear(128,2))\n    def forward(self,x):\n        x=self.conv_seq(x)\n        x = torch.flatten(x,1)\n        output=self.fc_seq(x)\n        return output","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CNN().to(DEVICE)\nprint(_color.S+str(model)+_color.E)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Now we use summary of pytorch to see how our model work**","metadata":{}},{"cell_type":"code","source":"summary(model,input_size=(1, 227, 227))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"About 51 million Params hmm :v","metadata":{}},{"cell_type":"markdown","source":"#### loss,optimizer","metadata":{}},{"cell_type":"markdown","source":"According to the paper we has show on top we will use our loss function is Crossentropy,\nand optimizer is RMS prop\n![image2](https://scontent.fsgn13-2.fna.fbcdn.net/v/t1.15752-9/325374357_555040573207258_5656297663407678770_n.png?_nc_cat=105&ccb=1-7&_nc_sid=ae9488&_nc_ohc=dz_kGVtlOIIAX-27NL-&_nc_ht=scontent.fsgn13-2.fna&oh=03_AdS_TWm668YiqqB4RK5xHNhjq83dvdlM0BGix5O0hH_oDw&oe=63EAE402)","metadata":{}},{"cell_type":"code","source":"loss_fn = nn.CrossEntropyLoss()\noptimizer = torch.optim.RMSprop(model.parameters(), lr=1e-3)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### train loop","metadata":{}},{"cell_type":"markdown","source":"Train function form [Pytorch Quickstart](https://pytorch.org/tutorials/beginner/basics/quickstart_tutorial.html)","metadata":{}},{"cell_type":"code","source":"def train(train_dataloader, model, loss_fn, optimizer):\n    size = len(train_dataloader.dataset)\n    model.train()\n    for batch, (img, la) in enumerate(train_dataloader):\n        img, la = data_to_device(img, la)\n        \n        # Compute prediction error\n        pred = model(img)\n        loss = loss_fn(pred, la)\n\n        # Backpropagation\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        wandb.log({'Lost of model':loss})\n        if batch % 1000 == 0:\n            loss, current = loss.item(), batch * len(img)\n            print(f\"loss: {loss:>7f}  [{current:>5d}/{size:>5d}]\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Let's Train 2 epoch","metadata":{}},{"cell_type":"code","source":"epochs = 2\nrun = wandb.init(project=\"RSNA_train\",reinit=True)\nwandb.run.name = \"RSNA do something with loss func\"\nfor t in range(epochs):\n    print(f\"Epoch {t+1}\\n-------------------------------\")\n    train(train_dataloader, model, loss_fn, optimizer)\nprint(\"Done!\")\nrun.finish()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Data imbalace so that loss score seem crazy lol**","metadata":{}},{"cell_type":"markdown","source":"#### Predict","metadata":{}},{"cell_type":"code","source":"classes = ['No Cancer','Cancer']\nlist_pred_id=[]\nlist_pred_cancer=[]\nmodel.eval()\nwith torch.no_grad():\n    fig,ax = plt.subplots(1,4,figsize=(15,15))\n    for k,(img_test,pred_id) in enumerate(test_dataset):\n        ax_idx = ax[k]\n        ax_idx.imshow(img_test.permute(1,2,0),cmap=plt.cm.gray)\n        pred = model(img_test.type(torch.cuda.FloatTensor).unsqueeze(0))\n        softmax=nn.Softmax(dim=1)\n        final_pred = softmax(pred)\n        predicted = classes[final_pred[0].argmax(0)]\n        list_pred_id.append(pred_id)\n        list_pred_cancer.append(final_pred[0].argmax(0).item())\n        ax_idx.set_title(f\"Fig {pred_id} is {predicted}\")\n\n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"list_pred_id,list_pred_cancer","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission = pd.DataFrame({'prediction_id':list_pred_id,\n                              'cancer':list_pred_cancer})\ndf_submission.to_csv('submission.csv',index=False)\npd.read_csv('submission.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class LightningModel(pl.LightningModule):\n#     def __init__(self, batch_size = 32, lr=1e-3):\n#         super().__init__()\n#         def layer(input_channel: int, \n#                   output_channel: int, \n#                   kernel_size_conv: int=3, \n#                   padding: str='same', \n#                   stride_conv: int=1, \n#                   kernel_size_maxpool: int=2, \n#                   stride_maxpool: int=2,\n#                   batchnorm=True,\n#                   maxpool=True,):\n#             layers = [nn.Conv2d(input_channel, output_channel, kernel_size_conv,stride_conv, padding)]\n#             if batchnorm:\n#                 layers.append(nn.BatchNorm2d(output_channel))\n                \n#             layers.append(nn.ReLU())\n            \n#             if maxpool:\n#                 layers.append(nn.MaxPool2d(kernel_size_maxpool, stride_maxpool))\n                \n#             return layers\n        \n#         self.conv_seq = nn.Sequential(\n#                         *layer(input_channel=1, \n#                                output_channel=8 ,\n#                                kernel_size_conv=3, \n#                                stride_conv=1,\n#                                padding='same',\n#                                kernel_size_maxpool=2,\n#                                stride_maxpool=2),\n            \n#                         *layer(input_channel=8, \n#                                output_channel=16 ,\n#                                kernel_size_conv=3, \n#                                stride_conv=1,\n#                                padding='same',\n#                                kernel_size_maxpool=2,\n#                                stride_maxpool=2),\n            \n#                         *layer(input_channel=16, \n#                                output_channel=32,\n#                                kernel_size_conv=3, \n#                                stride_conv=1,\n#                                padding='same',\n#                                maxpool=False),\n#                     )\n#         self.fc_seq = nn.Sequential(nn.Linear(100352,512),\n#                                    nn.ReLU(),\n#                                    nn.Linear(512,128),\n#                                    nn.ReLU(),\n#                                    nn.Linear(128,2))\n#     def forward(self,x):\n#         x=self.conv_seq(x)\n#         x = torch.flatten(x,1)\n#         output=self.fc_seq(x)\n#         return output\n    \n#     def training_step(self,batch,batch_idx):\n#         img,label = batch\n#         img=img.view(img.size(0),-1)\n#         pred_img = self.forward()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}