{"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 torchvision.transforms import ToTensor\nfrom torch.utils.data import TensorDataset, ConcatDataset, DataLoader, Dataset\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,\n                            Resize,  ColorJitter, GaussianBlur, RandomBrightnessContrast)\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\nfrom PIL import Image\nfrom sklearn.metrics import confusion_matrix","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-10T17:34:50.893568Z","iopub.execute_input":"2023-05-10T17:34:50.894193Z","iopub.status.idle":"2023-05-10T17:34:55.501891Z","shell.execute_reply.started":"2023-05-10T17:34:50.893824Z","shell.execute_reply":"2023-05-10T17:34:55.500682Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pip install wandb","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:34:55.504306Z","iopub.execute_input":"2023-05-10T17:34:55.504958Z","iopub.status.idle":"2023-05-10T17:34:55.511465Z","shell.execute_reply.started":"2023-05-10T17:34:55.504918Z","shell.execute_reply":"2023-05-10T17:34:55.508945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# # pass: 5ea6bd91c3e49f50e2842e8fc29f928eb0f5cd82\n# import wandb\n# !wandb login 5ea6bd91c3e49f50e2842e8fc29f928eb0f5cd82\n# wandb.init(project=\"breastCancer-restnest-test\", entity=\"breast-cancer-kltn\")","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:34:55.513901Z","iopub.execute_input":"2023-05-10T17:34:55.514341Z","iopub.status.idle":"2023-05-10T17:34:55.523939Z","shell.execute_reply.started":"2023-05-10T17:34:55.514303Z","shell.execute_reply":"2023-05-10T17:34:55.522329Z"},"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-05-10T17:34:55.526851Z","iopub.execute_input":"2023-05-10T17:34:55.528864Z","iopub.status.idle":"2023-05-10T17:34:55.624727Z","shell.execute_reply.started":"2023-05-10T17:34:55.528826Z","shell.execute_reply":"2023-05-10T17:34:55.619605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dataset(Dataset):\n    def __init__(self, root_dir, transform=None):\n        self.root_dir = root_dir\n        self.dataframe = pd.read_csv(f'{root_dir}/description.csv')\n        self.dataframe = self.dataframe.sample(frac = 1)\n        self.transform = transform\n            \n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self,index):\n        image_path = f'{self.root_dir}/{self.dataframe.iloc[index].Path_save}'\n        image = cv2.imread(image_path, 0)\n        image = pre_img(image)\n        if self.transform != None:\n            image_trans = self.transform(image=image)['image']\n        else:\n            image_trans = image\n        label = (self.dataframe.iloc[index]['Cancer'] == 1)*1\n        return image_trans, label","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:34:55.630205Z","iopub.execute_input":"2023-05-10T17:34:55.631182Z","iopub.status.idle":"2023-05-10T17:34:55.642955Z","shell.execute_reply.started":"2023-05-10T17:34:55.631089Z","shell.execute_reply":"2023-05-10T17:34:55.641929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pre_img(img):\n    img = img.astype(np.uint8)\n    clahe = cv2.createCLAHE(clipLimit=4.0, tileGridSize=(8, 8))\n    img_clahe = clahe.apply(img)\n    return img_clahe","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:34:55.644711Z","iopub.execute_input":"2023-05-10T17:34:55.645512Z","iopub.status.idle":"2023-05-10T17:34:55.657371Z","shell.execute_reply.started":"2023-05-10T17:34:55.645434Z","shell.execute_reply":"2023-05-10T17:34:55.655334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = Compose([\n    Resize(height=512, width=512, always_apply=True),\n    Normalize(mean=0.449, std=0.226),\n    HorizontalFlip(),\n    VerticalFlip(),\n    Rotate(),\n#     RandomCrop(height=400, width=400),\n#     ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.2),\n    GaussianBlur(),\n#     RandomRotation(limit=30),\n#     RandomBrightnessContrast(),\n    ToTensorV2()\n])\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-05-10T17:34:55.660479Z","iopub.execute_input":"2023-05-10T17:34:55.661177Z","iopub.status.idle":"2023-05-10T17:34:55.670041Z","shell.execute_reply.started":"2023-05-10T17:34:55.661148Z","shell.execute_reply":"2023-05-10T17:34:55.669099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inbreast_img = '/kaggle/input/inbreast-roi-mammography'\n\nmias_img = '/kaggle/input/mias-roi-mammography'\n\nddsm_img = '/kaggle/input/mini-ddsm-roi-mammography'\n\nddsm = Dataset(ddsm_img, transform)\ntrain_ddsm = DataLoader(ddsm, batch_size=32, shuffle=True)\n\nmias = Dataset(mias_img, transform)\ntrain_mias = DataLoader(mias, batch_size=32, shuffle = True)\n\ninbreast = Dataset(inbreast_img, transform)\ntest_inbreast = DataLoader(inbreast, batch_size=32, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:34:55.671611Z","iopub.execute_input":"2023-05-10T17:34:55.672428Z","iopub.status.idle":"2023-05-10T17:34:55.746037Z","shell.execute_reply.started":"2023-05-10T17:34:55.672304Z","shell.execute_reply":"2023-05-10T17:34:55.744986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = ConcatDataset([ddsm, mias, inbreast])\nfull_dataset = DataLoader(data, batch_size=16, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:34:55.749245Z","iopub.execute_input":"2023-05-10T17:34:55.750139Z","iopub.status.idle":"2023-05-10T17:34:55.756592Z","shell.execute_reply.started":"2023-05-10T17:34:55.7501Z","shell.execute_reply":"2023-05-10T17:34:55.755206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader, random_split\ntrain_size = int(0.8 * len(data))\ntest_size = len(data) - train_size\n\n# Split the dataset into training and test sets\ntrain_dataset, test_dataset = random_split(data, [train_size, test_size])\n\n# Create data loaders for training and testing\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:34:55.76473Z","iopub.execute_input":"2023-05-10T17:34:55.765336Z","iopub.status.idle":"2023-05-10T17:34:55.774685Z","shell.execute_reply.started":"2023-05-10T17:34:55.765298Z","shell.execute_reply":"2023-05-10T17:34:55.773492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(train_loader)*32)\nprint(len(test_loader)*32)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:34:55.778965Z","iopub.execute_input":"2023-05-10T17:34:55.781213Z","iopub.status.idle":"2023-05-10T17:34:55.78885Z","shell.execute_reply.started":"2023-05-10T17:34:55.781178Z","shell.execute_reply":"2023-05-10T17:34:55.78788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for k,(img,la) in enumerate(train_loader):\n    if k == 2:\n        break\n    print(la)\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)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:34:55.790475Z","iopub.execute_input":"2023-05-10T17:34:55.791612Z","iopub.status.idle":"2023-05-10T17:35:04.346064Z","shell.execute_reply.started":"2023-05-10T17:34:55.791576Z","shell.execute_reply":"2023-05-10T17:35:04.344979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15, 15))\nfor i, (img, label) in enumerate(data):\n    plt.subplot(1,15,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 == 14:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:35:04.347872Z","iopub.execute_input":"2023-05-10T17:35:04.348262Z","iopub.status.idle":"2023-05-10T17:35:05.950223Z","shell.execute_reply.started":"2023-05-10T17:35:04.348221Z","shell.execute_reply":"2023-05-10T17:35:05.949254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test With RestNest","metadata":{}},{"cell_type":"code","source":"\nclass Bottleneck(nn.Module):\n    expansion = 4\n    def __init__(self, in_channels, out_channels, i_downsample=None, stride=1):\n        super(Bottleneck, self).__init__()\n        \n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)\n        self.batch_norm1 = nn.BatchNorm2d(out_channels)\n        \n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=stride, padding=1)\n        self.batch_norm2 = nn.BatchNorm2d(out_channels)\n        \n        self.conv3 = nn.Conv2d(out_channels, out_channels*self.expansion, kernel_size=1, stride=1, padding=0)\n        self.batch_norm3 = nn.BatchNorm2d(out_channels*self.expansion)\n        \n        self.i_downsample = i_downsample\n        self.stride = stride\n        self.relu = nn.ReLU()\n        \n    def forward(self, x):\n        identity = x.clone()\n        x = self.relu(self.batch_norm1(self.conv1(x)))\n        \n        x = self.relu(self.batch_norm2(self.conv2(x)))\n        \n        x = self.conv3(x)\n        x = self.batch_norm3(x)\n        \n        #downsample if needed\n        if self.i_downsample is not None:\n            identity = self.i_downsample(identity)\n        #add identity\n        x+=identity\n        x=self.relu(x)\n        \n        return x\n\nclass Block(nn.Module):\n    expansion = 1\n    def __init__(self, in_channels, out_channels, i_downsample=None, stride=1):\n        super(Block, self).__init__()\n       \n\n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1, stride=stride, bias=False)\n        self.batch_norm1 = nn.BatchNorm2d(out_channels)\n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1, stride=stride, bias=False)\n        self.batch_norm2 = nn.BatchNorm2d(out_channels)\n\n        self.i_downsample = i_downsample\n        self.stride = stride\n        self.relu = nn.ReLU()\n\n    def forward(self, x):\n        identity = x.clone()\n\n        x = self.relu(self.batch_norm2(self.conv1(x)))\n        x = self.batch_norm2(self.conv2(x))\n\n        if self.i_downsample is not None:\n            identity = self.i_downsample(identity)\n        print(x.shape)\n        print(identity.shape)\n        x += identity\n        x = self.relu(x)\n        return x\n\n\n        \n        \nclass ResNet(nn.Module):\n    def __init__(self, ResBlock, layer_list, num_classes, num_channels=1):\n        super(ResNet, self).__init__()\n        self.in_channels = 64\n        \n        self.conv1 = nn.Conv2d(num_channels, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        self.batch_norm1 = nn.BatchNorm2d(64)\n        self.relu = nn.ReLU()\n        self.max_pool = nn.MaxPool2d(kernel_size = 3, stride=2, padding=1)\n        \n        self.layer1 = self._make_layer(ResBlock, layer_list[0], planes=64)\n        self.layer2 = self._make_layer(ResBlock, layer_list[1], planes=128, stride=2)\n        self.layer3 = self._make_layer(ResBlock, layer_list[2], planes=256, stride=2)\n        self.layer4 = self._make_layer(ResBlock, layer_list[3], planes=512, stride=2)\n        \n        self.avgpool = nn.AdaptiveAvgPool2d((1,1))\n        self.fc = nn.Linear(512*ResBlock.expansion, num_classes)\n        \n    def forward(self, x):\n        x = self.relu(self.batch_norm1(self.conv1(x)))\n        x = self.max_pool(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.reshape(x.shape[0], -1)\n        x = self.fc(x)\n        \n        return x\n        \n    def _make_layer(self, ResBlock, blocks, planes, stride=1):\n        ii_downsample = None\n        layers = []\n        \n        if stride != 1 or self.in_channels != planes*ResBlock.expansion:\n            ii_downsample = nn.Sequential(\n                nn.Conv2d(self.in_channels, planes*ResBlock.expansion, kernel_size=1, stride=stride),\n                nn.BatchNorm2d(planes*ResBlock.expansion)\n            )\n            \n        layers.append(ResBlock(self.in_channels, planes, i_downsample=ii_downsample, stride=stride))\n        self.in_channels = planes*ResBlock.expansion\n        \n        for i in range(blocks-1):\n            layers.append(ResBlock(self.in_channels, planes))\n            \n        return nn.Sequential(*layers)\n\n        \n        \ndef ResNet50(num_classes, channels=1):\n    return ResNet(Bottleneck, [3,4,6,3], num_classes, channels)\n    \ndef ResNet101(num_classes, channels=1):\n    return ResNet(Bottleneck, [3,4,23,3], num_classes, channels)\n\ndef ResNet152(num_classes, channels=1):\n    return ResNet(Bottleneck, [3,8,36,3], num_classes, channels)\n","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:35:05.951808Z","iopub.execute_input":"2023-05-10T17:35:05.952482Z","iopub.status.idle":"2023-05-10T17:35:05.985069Z","shell.execute_reply.started":"2023-05-10T17:35:05.952439Z","shell.execute_reply":"2023-05-10T17:35:05.983988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"net = ResNet152(2).to(DEVICE)\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.SGD(net.parameters(), lr=0.1, momentum=0.9, weight_decay=0.0001)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor = 0.1, patience=5)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:35:05.986748Z","iopub.execute_input":"2023-05-10T17:35:05.987125Z","iopub.status.idle":"2023-05-10T17:35:06.681334Z","shell.execute_reply.started":"2023-05-10T17:35:05.987081Z","shell.execute_reply":"2023-05-10T17:35:06.680329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working/model_cnn_breast_cancer","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:35:06.682758Z","iopub.execute_input":"2023-05-10T17:35:06.683218Z","iopub.status.idle":"2023-05-10T17:35:07.665137Z","shell.execute_reply.started":"2023-05-10T17:35:06.683179Z","shell.execute_reply":"2023-05-10T17:35:07.663773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS = 20\nfor epoch in range(EPOCHS):\n    losses = []\n    running_loss = 0\n    running_corrects = 0\n    for i, inp in enumerate(train_loader):\n        inputs, labels = inp\n        inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n        optimizer.zero_grad()\n\n        outputs = net(inputs)\n        loss = criterion(outputs, labels)\n        losses.append(loss.item())\n\n        _, preds = torch.max(outputs, 1)\n        batch_corrects = torch.sum(preds == labels.data)\n        running_corrects += batch_corrects\n\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n\n        if i%100 == 0 and i > 0:\n            batch_acc = batch_corrects.double() / 16\n            print(f'Loss [{epoch+1}, {i}](epoch, minibatch): ', running_loss / 100, \n                  '\\n \\t Batch accuracy: ', batch_acc.item())\n            wandb.log({'loss':running_loss/100, \n                       'batch_accuracy':batch_acc.item()})\n            running_loss = 0.0\n\n        torch.save(net.state_dict(), \"/kaggle/working/model_cnn_breast_cancer/resnet152.pth\")\n\n    avg_loss = sum(losses)/len(losses)\n    avg_acc = running_corrects.double() / (len(train_dataloader.dataset))\n    scheduler.step(avg_loss)\n\n    print(f'Epoch [{epoch+1}] average loss: {avg_loss}, average accuracy: {avg_acc}')\n    wandb.log({'epoch_average_loss':avg_loss, 'epoch_average_accuracy':avg_acc})\n    \nprint('Training Done')","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:35:07.667615Z","iopub.execute_input":"2023-05-10T17:35:07.668828Z","iopub.status.idle":"2023-05-10T17:35:15.037792Z","shell.execute_reply.started":"2023-05-10T17:35:07.668778Z","shell.execute_reply":"2023-05-10T17:35:15.036111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:35:15.038928Z","iopub.status.idle":"2023-05-10T17:35:15.040672Z","shell.execute_reply.started":"2023-05-10T17:35:15.040399Z","shell.execute_reply":"2023-05-10T17:35:15.040424Z"},"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(test_dataset):\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-05-10T17:35:15.042295Z","iopub.status.idle":"2023-05-10T17:35:15.04309Z","shell.execute_reply.started":"2023-05-10T17:35:15.042826Z","shell.execute_reply":"2023-05-10T17:35:15.04285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classes = ['No Cancer','Cancer']\nlist_pred_id=[]\nlist_pred_cancer=[]\nnet.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_loader):\n        ax_idx = ax[k]\n        ax_idx.imshow(img_test.permute(1,2,0),cmap=plt.cm.gray)\n        pred = net(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}\")","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:35:15.044562Z","iopub.status.idle":"2023-05-10T17:35:15.045372Z","shell.execute_reply.started":"2023-05-10T17:35:15.045083Z","shell.execute_reply":"2023-05-10T17:35:15.045119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\ncorrect = 0\ntotal = 0\ntrue_labels = []\npredicted_labels = []\n\nwith torch.no_grad():\n    for data in train_ddsm:\n        images, labels = data\n        images = images.to(DEVICE)\n        labels = labels.to(DEVICE)\n\n        outputs = model(images)\n        probabilities = F.softmax(outputs, dim=1)\n        predicted = torch.argmax(probabilities, 1)\n\n        true_labels.extend(labels.cpu().numpy())\n        predicted_labels.extend(predicted.cpu().numpy())\n        \n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n        \n#         print('probabilities: ', probabilities)\n#         print('correct: ', correct)\n#         print('predicted: ', predicted)\n#         print('labels: ', labels)\n\naccuracy = 100 * correct / total\nprint(f\"Accuracy: {accuracy:.2f}%\")\n\n# Calculate the confusion matrix\ncm = confusion_matrix(true_labels, predicted_labels)\n\n# Print the confusion matrix\nprint(\"Confusion Matrix:\")\nprint(cm)","metadata":{"execution":{"iopub.status.busy":"2023-05-10T17:35:15.046801Z","iopub.status.idle":"2023-05-10T17:35:15.047605Z","shell.execute_reply.started":"2023-05-10T17:35:15.047341Z","shell.execute_reply":"2023-05-10T17:35:15.047366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\ncorrect = 0\ntotal = 0\ntrue_labels = []\npredicted_labels = []\n\nwith torch.no_grad():\n    for data in train_mias:\n        images, labels = data\n        images = images.to(DEVICE)\n        labels = labels.to(DEVICE)\n\n        outputs = model(images)\n        probabilities = F.softmax(outputs, dim=1)\n        predicted = torch.argmax(probabilities, 1)\n\n        true_labels.extend(labels.cpu().numpy())\n        predicted_labels.extend(predicted.cpu().numpy())\n        \n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n        \n#         print('probabilities: ', probabilities)\n#         print('correct: ', correct)\n#         print('predicted: ', predicted)\n#         print('labels: ', labels)\n\naccuracy = 100 * correct / total\nprint(f\"Accuracy: {accuracy:.2f}%\")\n\n# Calculate the confusion matrix\ncm = confusion_matrix(true_labels, predicted_labels)\n\n# Print the confusion matrix\nprint(\"Confusion Matrix:\")\nprint(cm)","metadata":{},"execution_count":null,"outputs":[]}]}