{"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 pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport torch\nfrom torchvision import datasets\nimport torchvision.transforms as transforms\nfrom torchvision.io import read_image\nfrom torch.utils.data import Dataset\nfrom torchvision.transforms import ToTensor\nfrom torch.utils.data import DataLoader\nimport os\nimport cv2\nfrom skimage import io\nfrom skimage import data\nfrom skimage import filters\nimport glob, itertools\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\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport pydicom\nimport scipy\nfrom skimage.exposure import equalize_adapthist\nimport tqdm\nimport logging\nimport torch.optim as optim","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-03-26T07:16:28.58115Z","iopub.execute_input":"2023-03-26T07:16:28.581564Z","iopub.status.idle":"2023-03-26T07:16:33.934495Z","shell.execute_reply.started":"2023-03-26T07:16:28.581535Z","shell.execute_reply":"2023-03-26T07:16:33.933023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class _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":{"execution":{"iopub.status.busy":"2023-03-26T07:16:33.936591Z","iopub.execute_input":"2023-03-26T07:16:33.937314Z","iopub.status.idle":"2023-03-26T07:16:33.990556Z","shell.execute_reply.started":"2023-03-26T07:16:33.937277Z","shell.execute_reply":"2023-03-26T07:16:33.98699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\ndef 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":{"execution":{"iopub.status.busy":"2023-03-26T07:16:33.994051Z","iopub.execute_input":"2023-03-26T07:16:33.994338Z","iopub.status.idle":"2023-03-26T07:16:34.006816Z","shell.execute_reply.started":"2023-03-26T07:16:33.994312Z","shell.execute_reply":"2023-03-26T07:16:34.005673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DDSM_dataset(Dataset):\n    def __init__(self,excel_file,root_dir, lst, option, transform=None, is_train = True):\n        \"\"\"\n        is_train: if it's DDSM dataset or not it's RSNA data set\n        \"\"\"\n        self.root_dir = root_dir\n        self.is_train = is_train\n        if self.is_train:\n            self.dataframe = pd.read_excel(excel_file)\n            self.dataframe = self.dataframe[(self.dataframe['Status']!='Benign') | (self.dataframe['fileName'].isin(lst))]\n            if option == 'CC_L':\n                self.dataframe = self.dataframe[(self.dataframe['View']=='CC') & (self.dataframe['Side']=='LEFT')]\n            elif option == 'CC_R':\n                self.dataframe = self.dataframe[(self.dataframe['View']=='CC') & (self.dataframe['Side']=='RIGHT')]\n            elif option == 'MLO_L':\n                self.dataframe = self.dataframe[(self.dataframe['View']=='MLO') & (self.dataframe['Side']=='LEFT')]\n            else:\n                self.dataframe = self.dataframe[(self.dataframe['View']=='MLO') & (self.dataframe['Side']=='RIGHT')]\n            self.dataframe = self.dataframe.sample(frac = 1)\n            self.transform = transform\n\n        else:\n            self.dataframe = pd.read_csv(excel_file)\n            self.transform = Compose([Resize(height=227,width=227,always_apply=True),\n                                      ToTensorV2()])\n            if option == 'CC_L':\n                self.dataframe = self.dataframe[(self.dataframe['view']=='CC') & (self.dataframe['laterality']=='L')]\n            elif option == 'CC_R':\n                self.dataframe = self.dataframe[(self.dataframe['view']=='CC') & (self.dataframe['laterality']=='R')]\n            elif option == 'MLO_L':\n                self.dataframe = self.dataframe[(self.dataframe['view']=='MLO') & (self.dataframe['laterality']=='L')]\n            else:\n                self.dataframe = self.dataframe[(self.dataframe['view']=='MLO') & (self.dataframe['laterality']=='R')]\n            \n        \n    def __len__(self):\n        return len(self.dataframe)\n    \n    def pre_img(self, img):\n        clahe = cv2.createCLAHE(clipLimit=4.0, tileGridSize=(8, 8))\n        img_clahe = clahe.apply(img)\n        return img_clahe\n    \n    def __getitem__(self,index):\n        \n        if self.is_train:\n            if (self.dataframe.iloc[index]['Status']=='Normal') :\n                image_path = self.root_dir + 'Normal/' + str(self.dataframe.iloc[index].fullPath.replace(\"\\\\\", \"/\").split('/')[2])\n            elif (self.dataframe.iloc[index]['Status']=='Cancer') :\n                image_path = self.root_dir + 'Cancer/' + str(self.dataframe.iloc[index].fullPath.replace(\"\\\\\", \"/\").split('/')[2])\n            image = cv2.imread(image_path,0)\n            image = self.pre_img(image)\n            if self.transform != None:\n                image_trans = self.transform(image=image)['image']\n\n            else:\n                image_trans = image \n            label = (self.dataframe.iloc[index]['Status'] == \"Cancer\")*1\n            return image_trans,label\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            if self.transform != None:\n                image_trans = self.transform(image=image)['image']\n            prediction_id = self.dataframe.iloc[index]['prediction_id']\n            return image_trans,prediction_id","metadata":{"execution":{"iopub.status.busy":"2023-03-26T07:16:34.010031Z","iopub.execute_input":"2023-03-26T07:16:34.01087Z","iopub.status.idle":"2023-03-26T07:16:34.02963Z","shell.execute_reply.started":"2023-03-26T07:16:34.010832Z","shell.execute_reply":"2023-03-26T07:16:34.028497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"excel_path = '/kaggle/input/miniddsm2/MINI-DDSM-Complete-JPEG-8/DataWMask.xlsx'\nimage_path = '/kaggle/input/ddsm-croped-image/'\ntest_csv_path = '/kaggle/input/rsna-breast-cancer-detection/test.csv'\ntest_RSNA = '/kaggle/input/rsna-breast-cancer-detection/test_images'\nlst = ['A_1000_1.LEFT_MLO.jpg', 'A_1000_1.RIGHT_MLO.jpg', 'A_1583_1.RIGHT_MLO.jpg','B_3476_1.RIGHT_CC.jpg',\n                  'B_3497_1.LEFT_CC.jpg', 'B_3497_1.LEFT_MLO.jpg', 'B_3497_1.RIGHT_CC.jpg', 'A_1589_1.RIGHT_MLO.jpg',\n                  'B_3084_1.LEFT_MLO.jpg', 'B_3474_1.LEFT_CC.jpg', 'B_3474_1.RIGHT_CC.jpg', 'B_3476_1.RIGHT_CC.jpg',\n                  'B_3497_1.RIGHT_CC.jpg', 'B_3497_1.RIGHT_MLO.jpg', 'B_3500_1.LEFT_CC.jpg']\ntransform = Compose([Resize(height=227,width=227,always_apply=True),\n                    Normalize(mean=0.449,std=0.226),\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":{"execution":{"iopub.status.busy":"2023-03-26T07:16:34.031202Z","iopub.execute_input":"2023-03-26T07:16:34.031555Z","iopub.status.idle":"2023-03-26T07:16:34.04224Z","shell.execute_reply.started":"2023-03-26T07:16:34.031519Z","shell.execute_reply":"2023-03-26T07:16:34.041129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CC_L=DDSM_dataset(excel_path,image_path, lst, \"CC_L\", transform)\ntrain_CC_L = DataLoader(CC_L, batch_size=32, shuffle=True)\nCC_R=DDSM_dataset(excel_path,image_path, lst, \"CC_R\", transform)\ntrain_CC_R = DataLoader(CC_R, batch_size=32, shuffle=True)\nMLO_L=DDSM_dataset(excel_path,image_path, lst, \"MLO_L\", transform)\ntrain_MLO_L = DataLoader(MLO_L, batch_size=32, shuffle=True)\nMLO_R=DDSM_dataset(excel_path,image_path, lst, \"MLO_R\", transform)\ntrain_MLO_R = DataLoader(MLO_R, batch_size=32, shuffle=True)\n\ntest_CC_L_data=DDSM_dataset(test_csv_path,test_RSNA, lst, 'CC_L', is_train = False)\ntest_CC_L = DataLoader(test_CC_L_data, batch_size=1, shuffle=False)\ntest_CC_R_data=DDSM_dataset(test_csv_path,test_RSNA, lst, 'CC_R', is_train = False)\ntest_CC_L = DataLoader(test_CC_R_data, batch_size=1, shuffle=False)\ntest_MLO_L_data=DDSM_dataset(test_csv_path,test_RSNA, lst, 'MLO_L', is_train = False)\ntest_MLO_L = DataLoader(test_MLO_L_data, batch_size=1, shuffle=False)\ntest_MLO_R_data=DDSM_dataset(test_csv_path,test_RSNA, lst, 'MLO_R', is_train = False)\ntest_MLP_R = DataLoader(test_MLO_R_data, batch_size=1, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-03-26T07:16:34.044615Z","iopub.execute_input":"2023-03-26T07:16:34.045672Z","iopub.status.idle":"2023-03-26T07:16:40.50195Z","shell.execute_reply.started":"2023-03-26T07:16:34.045635Z","shell.execute_reply":"2023-03-26T07:16:40.500942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15, 15))\nfor i, (img, label) in enumerate(CC_L):\n    plt.subplot(1,5,i+1)\n    plt.imshow(img.squeeze(), cmap='gray')\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    plt.title(label)\n    if i == 4:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:12:34.028622Z","iopub.execute_input":"2023-03-25T17:12:34.029782Z","iopub.status.idle":"2023-03-25T17:12:34.488332Z","shell.execute_reply.started":"2023-03-25T17:12:34.029737Z","shell.execute_reply":"2023-03-25T17:12:34.487258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_CC_L.__len__()*32, train_CC_R.__len__()*32, train_MLO_L.__len__()*32, train_MLO_R.__len__()*32","metadata":{"execution":{"iopub.status.busy":"2023-03-26T07:24:20.001028Z","iopub.execute_input":"2023-03-26T07:24:20.001457Z","iopub.status.idle":"2023-03-26T07:24:20.017025Z","shell.execute_reply.started":"2023-03-26T07:24:20.001422Z","shell.execute_reply":"2023-03-26T07:24:20.009667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df=pd.read_excel('/kaggle/input/miniddsm2/MINI-DDSM-Complete-JPEG-8/DataWMask.xlsx')\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2023-03-26T07:18:29.238686Z","iopub.execute_input":"2023-03-26T07:18:29.239609Z","iopub.status.idle":"2023-03-26T07:18:30.521397Z","shell.execute_reply.started":"2023-03-26T07:18:29.239571Z","shell.execute_reply":"2023-03-26T07:18:30.520259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = df[df['Status']!='Benign']\ndf.groupby(['Side'])['View'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2023-03-26T07:21:38.324914Z","iopub.execute_input":"2023-03-26T07:21:38.32586Z","iopub.status.idle":"2023-03-26T07:21:38.342899Z","shell.execute_reply.started":"2023-03-26T07:21:38.32581Z","shell.execute_reply":"2023-03-26T07:21:38.341885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15, 15))\nfor i, (img, label) in enumerate(CC_R):\n    plt.subplot(1,5,i+1)\n    plt.imshow(img.squeeze(), cmap='gray')\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    plt.title(label)\n    if i == 4:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:12:34.489754Z","iopub.execute_input":"2023-03-25T17:12:34.490719Z","iopub.status.idle":"2023-03-25T17:12:34.977772Z","shell.execute_reply.started":"2023-03-25T17:12:34.490679Z","shell.execute_reply":"2023-03-25T17:12:34.976897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15, 15))\nfor i, (img, label) in enumerate(CC_R):\n    plt.subplot(1,5,i+1)\n    plt.imshow(img.squeeze(), cmap='gray')\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    plt.title(label)\n    if i == 4:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:12:34.981467Z","iopub.execute_input":"2023-03-25T17:12:34.982521Z","iopub.status.idle":"2023-03-25T17:12:35.483427Z","shell.execute_reply.started":"2023-03-25T17:12:34.982476Z","shell.execute_reply":"2023-03-25T17:12:35.482373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15, 15))\nfor i, (img, label) in enumerate(MLO_L):\n    plt.subplot(1,5,i+1)\n    plt.imshow(img.squeeze(), cmap='gray')\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    plt.title(label)\n    if i == 4:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:12:35.486624Z","iopub.execute_input":"2023-03-25T17:12:35.486992Z","iopub.status.idle":"2023-03-25T17:12:35.965457Z","shell.execute_reply.started":"2023-03-25T17:12:35.486964Z","shell.execute_reply":"2023-03-25T17:12:35.964413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15, 15))\nfor i, (img, label) in enumerate(MLO_R):\n    plt.subplot(1,5,i+1)\n    plt.imshow(img.squeeze(), cmap='gray')\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    plt.title(label)\n    if i == 4:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:12:35.966777Z","iopub.execute_input":"2023-03-25T17:12:35.967614Z","iopub.status.idle":"2023-03-25T17:12:36.474703Z","shell.execute_reply.started":"2023-03-25T17:12:35.967584Z","shell.execute_reply":"2023-03-25T17:12:36.467245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\n\nclass Bottleneck(nn.Module):\n    expansion = 3  # Number of output channels of the block relative to the input channels\n\n    def __init__(self, in_channels, out_channels, stride=1, downsample=None, width=32):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_channels, width, kernel_size=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(width)\n        self.conv2 = nn.Conv2d(width, width, kernel_size=3, stride=stride, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(width)\n        self.conv3 = nn.Conv2d(width, out_channels * self.expansion, kernel_size=1, bias=False)\n        self.bn3 = nn.BatchNorm2d(out_channels * self.expansion)\n        self.relu = nn.ReLU(inplace=True)\n        self.downsample = downsample\n\n    def forward(self, x):\n        identity = 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            identity = self.downsample(x)\n\n        out += identity\n        out = self.relu(out)\n\n        return out\n\n\nclass HRNet(nn.Module):\n    def __init__(self, block, layers, num_classes=2, width=32):\n        super().__init__()\n        self.in_channels = 64\n\n        self.conv1 = nn.Conv2d(1, 64, kernel_size=3, stride=2, padding=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n\n        self.layer1 = self._make_layer(block, 64, layers[0], stride=1, width=width)\n        self.layer2 = self._make_layer(block, 128, layers[1], stride=2, width=width)\n        self.layer3 = self._make_layer(block, 256, layers[2], stride=2, width=width)\n        self.layer4 = self._make_layer(block, 512, layers[3], stride=2, width=width)\n\n        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))\n        self.fc = nn.Linear(512 * block.expansion, num_classes)\n\n    def _make_layer(self, block, out_channels, blocks, stride=1, width=32):\n        downsample = None\n        if stride != 1 or self.in_channels != out_channels * block.expansion:\n            downsample = nn.Sequential(\n                nn.Conv2d(self.in_channels, out_channels * block.expansion, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(out_channels * block.expansion),\n            )\n\n        layers = []\n        layers.append(block(self.in_channels, out_channels, stride, downsample, width=width))\n        self.in_channels = out_channels * block.expansion\n\n        for _ in range(1, blocks):\n            layers.append(block(self.in_channels, out_channels, width=width))\n\n        return nn.Sequential(*layers)\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\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        x = x.view(x.size(0), -1)\n        x = self.fc(x)\n\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:12:36.478829Z","iopub.execute_input":"2023-03-25T17:12:36.479183Z","iopub.status.idle":"2023-03-25T17:12:36.50402Z","shell.execute_reply.started":"2023-03-25T17:12:36.47914Z","shell.execute_reply":"2023-03-25T17:12:36.502965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working/model","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:12:36.508548Z","iopub.execute_input":"2023-03-25T17:12:36.509366Z","iopub.status.idle":"2023-03-25T17:12:37.554072Z","shell.execute_reply.started":"2023-03-25T17:12:36.50932Z","shell.execute_reply":"2023-03-25T17:12:37.552742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install wandb\nimport wandb\n!wandb login 5ea6bd91c3e49f50e2842e8fc29f928eb0f5cd82\nwandb.init(project=\"HR-Net\", entity=\"breast-cancer-kltn\")","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:12:37.556504Z","iopub.execute_input":"2023-03-25T17:12:37.557391Z","iopub.status.idle":"2023-03-25T17:13:25.567984Z","shell.execute_reply.started":"2023-03-25T17:12:37.557338Z","shell.execute_reply":"2023-03-25T17:13:25.566978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = HRNet(Bottleneck, [3, 4, 23, 3]).to(DEVICE)\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=0.01)","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:13:25.572165Z","iopub.execute_input":"2023-03-25T17:13:25.575791Z","iopub.status.idle":"2023-03-25T17:13:25.720498Z","shell.execute_reply.started":"2023-03-25T17:13:25.57575Z","shell.execute_reply":"2023-03-25T17:13:25.719445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import classification_report\n# Step 3: Train your model\nfor epoch in range(100):\n    running_loss = 0.0\n    running_corrects = 0\n    total_samples = 0\n    all_preds = []\n    all_labels = []\n    for i, inp in enumerate(train_MLO_L, 0):\n        inputs, labels = inp\n        inputs, labels = data_to_device(inputs, labels)\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        _, preds = torch.max(outputs, 1)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item()\n        running_corrects += torch.sum(preds == labels.data)\n        total_samples += len(labels)\n        all_preds.extend(preds.tolist())\n        all_labels.extend(labels.tolist())\n        if i % 10 == 5:\n            train_acc = running_corrects / total_samples\n            train_loss = running_loss / 100\n            print(f\"[Epoch {epoch + 1}, Batch {i }] loss: {train_loss:.3f}, acc: {train_acc:.3f}\")\n            wandb.log({'loss':train_loss, 'accuracy':train_acc})\n            running_loss = 0.0\n            running_corrects = 0\n            total_samples = 0\n\n    # Calculate F1-score, recall, and precision\n    report = classification_report(all_labels, all_preds, output_dict=True)\n    f1_score = report['weighted avg']['f1-score']\n    recall = report['weighted avg']['recall']\n    precision = report['weighted avg']['precision']\n    wandb.log({'f1-score':f1_score, 'recall':recall, 'precision':precision})\n\n# Step 5: Save your trained model\ntorch.save(model.state_dict(), \"/kaggle/working/model/MLO_L.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:16:24.126911Z","iopub.execute_input":"2023-03-25T17:16:24.127314Z","iopub.status.idle":"2023-03-25T17:18:08.415139Z","shell.execute_reply.started":"2023-03-25T17:16:24.127278Z","shell.execute_reply":"2023-03-25T17:18:08.413688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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,6,figsize=(15,15))\n    for k,(img_test,pred_id) in enumerate(MLO_R):\n        ax_idx = ax[k]\n        ax_idx.imshow(img_test.permute(1,2,0),cmap=plt.cm.gray)\n        print(len(img_test), pred_id)\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        if k == 5:\n            break","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:16:00.893692Z","iopub.status.idle":"2023-03-25T17:16:00.894463Z","shell.execute_reply.started":"2023-03-25T17:16:00.894202Z","shell.execute_reply":"2023-03-25T17:16:00.894229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Step 4: Evaluate your model\ncorrect = 0\ntotal = 0\nwith torch.no_grad():\n    for data in train_MLO_R:\n        images, labels = data\n        images, labels = images.type(torch.cuda.FloatTensor), labels.type(torch.cuda.FloatTensor)\n        outputs = model(images)\n        _, predicted = torch.max(outputs.data, 1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\nprint(f\"Accuracy on validation set: {correct / total * 100:.2f}%\")","metadata":{"execution":{"iopub.status.busy":"2023-03-25T17:16:00.895931Z","iopub.status.idle":"2023-03-25T17:16:00.896718Z","shell.execute_reply.started":"2023-03-25T17:16:00.896433Z","shell.execute_reply":"2023-03-25T17:16:00.896459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}