{"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":"!pip install dicomsdl\n!pip install timm\nimport numpy as np\nimport pandas as pd\nimport torch\nimport cv2\nimport os\nimport csv\nimport pydicom\nimport random\nimport dicomsdl\nimport time\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom timm import *\nfrom torch.utils.data import DataLoader, Dataset, random_split\nfrom torch import optim\nfrom torch.optim import lr_scheduler\nfrom glob import glob\nfrom PIL import Image, ImageFile\nimport albumentations\nfrom albumentations import *\nimport torchvision.transforms as transforms\nimport torchvision\nfrom tqdm import tqdm\nfrom torch.optim import Adam\nfrom torch.cuda import amp\nfrom sklearn import metrics\nfrom torch.optim.lr_scheduler import OneCycleLR, ReduceLROnPlateau\nimport matplotlib.pyplot as plt\ntorch.manual_seed(0)","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":27.159536,"end_time":"2022-12-29T10:37:54.442321","exception":false,"start_time":"2022-12-29T10:37:27.282785","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:42:56.080589Z","iopub.execute_input":"2022-12-31T00:42:56.081127Z","iopub.status.idle":"2022-12-31T00:43:25.525165Z","shell.execute_reply.started":"2022-12-31T00:42:56.081012Z","shell.execute_reply":"2022-12-31T00:43:25.524037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    training_set = '/kaggle/input/rsna-bcd-roi-1024x-png-dataset/train_images'\n    test_set = '/kaggle/input/rsna-breast-cancer-detection/test_images'\n    train_csv = '/kaggle/input/rsna-breast-cancer-detection/train.csv'\n    test_csv = '/kaggle/input/rsna-breast-cancer-detection/test.csv'\n    submission = '/kaggle/input/rsna-breast-cancer-detection/sample_submission.csv'\n    img_size = 512\n    batch_size = 24\n    lr = 1e-4\n    fold = 1\n    epochs = 3\n    loss = nn.BCEWithLogitsLoss()\n    n_accumulate = 4\n    T_max = 10\n    min_lr = 1e-6\n    weight_decay = 1e-6\n    temperature = 0.1\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"papermill":{"duration":0.893763,"end_time":"2022-12-29T10:37:55.347677","exception":false,"start_time":"2022-12-29T10:37:54.453914","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:43:53.369434Z","iopub.execute_input":"2022-12-31T00:43:53.370024Z","iopub.status.idle":"2022-12-31T00:43:53.436303Z","shell.execute_reply.started":"2022-12-31T00:43:53.369989Z","shell.execute_reply":"2022-12-31T00:43:53.43526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def img2tensor(img, dtype:np.dtype=np.float32):\n    if img.ndim==2: \n        img=np.expand_dims(img, 2)\n    img=np.transpose(img, (2, 0, 1))\n    return torch.from_numpy(img.astype(dtype, copy=False))\n\ndef img2roi(img, is_dicom=False):\n    \"\"\"\n    Returns ROI area in other words \n    cuts the image to a desired one\n    \n    Because there are machine label tags,\n    undesired details out of the breast image.\n    \"\"\"\n    if not is_dicom:\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n        \n    img = np.array(img * 255, dtype = np.uint8)\n    # Binarize the image\n    bin_img = cv2.threshold(img, 20, 255, cv2.THRESH_BINARY)[1]\n\n    # Make contours around the binarized image, keep only the largest contour\n    contours, _ = cv2.findContours(bin_img, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\n    contour = max(contours, key=cv2.contourArea)\n\n    # Find ROI from largest contour\n    ys = contour.squeeze()[:, 0]\n    xs = contour.squeeze()[:, 1]\n    roi =  img[np.min(xs):np.max(xs), np.min(ys):np.max(ys)]\n    \n    #print(f\"Shape of ROI image: {roi.shape}\")\n    \n    return roi\n\ndef make_fold(fold=0):\n    df = pd.read_csv('/kaggle/input/rsna-breast-cancer-detection/train.csv')\n    patient_id = df.patient_id.unique()\n    patient_id = sorted(patient_id)\n\n    num_fold=5\n    rs = np.random.RandomState(1234)\n    rs.shuffle(patient_id)\n    patient_id = np.array(patient_id)\n    f = np.arange(len(patient_id))%num_fold\n    train_id = patient_id[f!=fold]\n    valid_id = patient_id[f==fold]\n\n    train_df = df[df.patient_id.isin(train_id)].reset_index(drop=True)\n    valid_df = df[df.patient_id.isin(valid_id)].reset_index(drop=True)\n    return train_df, valid_df\n\ndef read_dicom(path, fix_monochrome = True):\n    dicom = dicomsdl.open(path)\n    data = dicom.pixelData(storedvalue=True)  # storedvalue = True for int16 return otherwise float32\n    data = data - np.min(data)\n    data = data / np.max(data)\n    if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = 1.0 - data\n    return data","metadata":{"papermill":{"duration":0.026367,"end_time":"2022-12-29T10:37:55.416679","exception":false,"start_time":"2022-12-29T10:37:55.390312","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:43:54.98415Z","iopub.execute_input":"2022-12-31T00:43:54.984764Z","iopub.status.idle":"2022-12-31T00:43:54.997599Z","shell.execute_reply.started":"2022-12-31T00:43:54.984729Z","shell.execute_reply":"2022-12-31T00:43:54.996545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainROIDataset(Dataset):\n    def __init__(self, training_set, train_csv, train, transform, fold):\n        self.training_set = training_set\n        self.train_csv = pd.read_csv(train_csv)\n        self.train_df, self.val_df = make_fold(fold)\n        self.train_df['path'] = CFG.training_set + '/' + self.train_df.patient_id.apply(lambda x: str(x)) + '/' + self.train_df.image_id.apply(lambda x: str(x)) + '.png'\n        self.val_df['path'] = CFG.training_set + '/' + self.val_df.patient_id.apply(lambda x: str(x)) + '/' + self.val_df.image_id.apply(lambda x: str(x)) + '.png'\n        self.train = train\n        self.transform = transform\n        self.imgs = {0:[], 1:[]}\n        self.train_paths = list(self.train_df.path)\n        self.val_paths = list(self.val_df.path)\n        self.classes = [0,1]\n        \n    def __len__(self):\n        return len(self.train_df) if self.train else len(self.val_df)\n    \n    def __getitem__(self, index):\n        class_name = random.choice(self.classes)\n        \n        if self.train:\n            index = index % len(self.train_df.loc[self.train_df.cancer == class_name])\n            image_path = self.train_df.loc[self.train_df.cancer == class_name].iloc[index].path\n        else:\n            index = index % len(self.val_df.loc[self.val_df.cancer == class_name])\n            image_path = self.val_df.loc[self.val_df.cancer == class_name].iloc[index].path\n        img = cv2.cvtColor(cv2.imread(image_path), cv2.COLOR_BGR2RGB)\n        #img = img2roi(img, is_dicom=False)\n        img = cv2.resize(img,(512,512))\n        if self.transform:\n            img = self.transform(img)\n            #img = augmented['image']\n        img = torch.tensor(img, dtype=torch.float)\n        target = torch.tensor(class_name, dtype=torch.float)\n        return img, target","metadata":{"papermill":{"duration":0.026757,"end_time":"2022-12-29T10:37:55.454352","exception":false,"start_time":"2022-12-29T10:37:55.427595","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:43:56.043182Z","iopub.execute_input":"2022-12-31T00:43:56.043947Z","iopub.status.idle":"2022-12-31T00:43:56.058018Z","shell.execute_reply.started":"2022-12-31T00:43:56.043911Z","shell.execute_reply":"2022-12-31T00:43:56.056869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_augmentation(p=1.0):\n    return Compose([\n        #Resize(CFG_SEGMENTATION.IMG_SIZE,CFG_SEGMENTATION.IMG_SIZE),\n        HorizontalFlip(),\n        VerticalFlip(),\n        RandomRotate90(),\n        ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.1, rotate_limit=15, p=0.9, border_mode=cv2.BORDER_REFLECT),\n        OneOf([\n            OpticalDistortion(p=0.3),\n            GridDistortion(p=.1),\n            PiecewiseAffine(p=0.3),\n        ], p=0.3),\n        OneOf([\n            CLAHE(clip_limit=2),\n            RandomBrightnessContrast(),\n        ], p=0.3)\n    ], p=p)","metadata":{"papermill":{"duration":0.01978,"end_time":"2022-12-29T10:37:55.484965","exception":false,"start_time":"2022-12-29T10:37:55.465185","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:43:57.08131Z","iopub.execute_input":"2022-12-31T00:43:57.08228Z","iopub.status.idle":"2022-12-31T00:43:57.089644Z","shell.execute_reply.started":"2022-12-31T00:43:57.082228Z","shell.execute_reply":"2022-12-31T00:43:57.088276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.2),\n            transforms.GaussianBlur(kernel_size=3, sigma=(0.1, 2.0)),\n            transforms.RandomVerticalFlip(),\n            transforms.RandomEqualize(),\n            transforms.ToTensor()\n        ])","metadata":{"papermill":{"duration":0.019077,"end_time":"2022-12-29T10:37:55.51484","exception":false,"start_time":"2022-12-29T10:37:55.495763","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:43:57.9328Z","iopub.execute_input":"2022-12-31T00:43:57.933164Z","iopub.status.idle":"2022-12-31T00:43:57.941925Z","shell.execute_reply.started":"2022-12-31T00:43:57.933127Z","shell.execute_reply":"2022-12-31T00:43:57.939148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TrainROIDataset(CFG.training_set, CFG.train_csv, True, transform, CFG.fold)","metadata":{"papermill":{"duration":0.262682,"end_time":"2022-12-29T10:37:55.845534","exception":false,"start_time":"2022-12-29T10:37:55.582852","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:43:59.576377Z","iopub.execute_input":"2022-12-31T00:43:59.577065Z","iopub.status.idle":"2022-12-31T00:43:59.814681Z","shell.execute_reply.started":"2022-12-31T00:43:59.577029Z","shell.execute_reply":"2022-12-31T00:43:59.813693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_dataset)","metadata":{"papermill":{"duration":0.021346,"end_time":"2022-12-29T10:37:55.878211","exception":false,"start_time":"2022-12-29T10:37:55.856865","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:44:00.471603Z","iopub.execute_input":"2022-12-31T00:44:00.472669Z","iopub.status.idle":"2022-12-31T00:44:00.480351Z","shell.execute_reply.started":"2022-12-31T00:44:00.472619Z","shell.execute_reply":"2022-12-31T00:44:00.479141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_dataset = TrainROIDataset(CFG.training_set, CFG.train_csv, False, transform, CFG.fold)","metadata":{"papermill":{"duration":0.21166,"end_time":"2022-12-29T10:37:56.129523","exception":false,"start_time":"2022-12-29T10:37:55.917863","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:44:01.501274Z","iopub.execute_input":"2022-12-31T00:44:01.502168Z","iopub.status.idle":"2022-12-31T00:44:01.680495Z","shell.execute_reply.started":"2022-12-31T00:44:01.50212Z","shell.execute_reply":"2022-12-31T00:44:01.679497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(val_dataset)","metadata":{"papermill":{"duration":0.021312,"end_time":"2022-12-29T10:37:56.162209","exception":false,"start_time":"2022-12-29T10:37:56.140897","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:44:01.96973Z","iopub.execute_input":"2022-12-31T00:44:01.97008Z","iopub.status.idle":"2022-12-31T00:44:01.976735Z","shell.execute_reply.started":"2022-12-31T00:44:01.970052Z","shell.execute_reply":"2022-12-31T00:44:01.97579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=CFG.batch_size, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=CFG.batch_size, shuffle=True, num_workers=2)","metadata":{"papermill":{"duration":0.021897,"end_time":"2022-12-29T10:37:56.282474","exception":false,"start_time":"2022-12-29T10:37:56.260577","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:44:03.839831Z","iopub.execute_input":"2022-12-31T00:44:03.840735Z","iopub.status.idle":"2022-12-31T00:44:03.847707Z","shell.execute_reply.started":"2022-12-31T00:44:03.840698Z","shell.execute_reply":"2022-12-31T00:44:03.846735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Model(torch.nn.Module):\n    def __init__(self, model_type='efficientnet_b4', pretrained=True, dropout=0.):\n        super().__init__()\n        self.model = create_model(model_type, pretrained=pretrained, num_classes=0, drop_rate=dropout)\n\n        self.backbone_dim = self.model(torch.randn(1, 3, 512, 512)).shape[-1]\n\n        self.nn_cancer = torch.nn.Sequential(\n            torch.nn.Linear(self.backbone_dim, 1),\n        )\n\n    def forward(self, x):\n        x = self.model(x)\n        cancer = self.nn_cancer(x).squeeze()\n        return cancer\n\n    def predict(self, x):\n        cancer = self.forward(x)\n        return torch.sigmoid(cancer)","metadata":{"papermill":{"duration":0.021322,"end_time":"2022-12-29T10:37:56.314818","exception":false,"start_time":"2022-12-29T10:37:56.293496","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:44:06.072073Z","iopub.execute_input":"2022-12-31T00:44:06.072493Z","iopub.status.idle":"2022-12-31T00:44:06.085879Z","shell.execute_reply.started":"2022-12-31T00:44:06.072456Z","shell.execute_reply":"2022-12-31T00:44:06.084335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_n_params(model):\n    pp=0\n    for p in list(model.parameters()):\n        nn=1\n        for s in list(p.size()):\n            nn = nn*s\n        pp += nn\n    return pp","metadata":{"papermill":{"duration":0.019117,"end_time":"2022-12-29T10:37:56.574129","exception":false,"start_time":"2022-12-29T10:37:56.555012","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:44:06.871173Z","iopub.execute_input":"2022-12-31T00:44:06.871673Z","iopub.status.idle":"2022-12-31T00:44:06.87771Z","shell.execute_reply.started":"2022-12-31T00:44:06.871641Z","shell.execute_reply":"2022-12-31T00:44:06.876599Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def probabilistic_f1(labels, predictions, beta=0.5):\n    y_true_count = 0\n    ctp = 0\n    cfp = 0\n\n    for idx in range(len(labels)):\n        prediction = min(max(predictions[idx], 0), 1)\n        if (labels[idx]):\n            y_true_count += 1\n            ctp += prediction\n            cfp += 1 - prediction\n        else:\n            cfp += prediction\n\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)\n        return result\n    else:\n        return 0","metadata":{"papermill":{"duration":0.021893,"end_time":"2022-12-29T10:37:56.70336","exception":false,"start_time":"2022-12-29T10:37:56.681467","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:44:07.575243Z","iopub.execute_input":"2022-12-31T00:44:07.576142Z","iopub.status.idle":"2022-12-31T00:44:07.583122Z","shell.execute_reply.started":"2022-12-31T00:44:07.5761Z","shell.execute_reply":"2022-12-31T00:44:07.582112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_loop(train_loader, val_loader):\n    n_epochs = CFG.epochs\n    train_loss = 0\n    scaler = torch.cuda.amp.GradScaler(enabled = True)\n    net = Model(model_type='efficientnet_b4',dropout=0.2).cuda()\n    optimizer = optim.Adam(net.parameters(), lr=1e-4, weight_decay=1e-2)\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=0.01, steps_per_epoch=len(train_loader), epochs=n_epochs)    \n    train_loss, val_loss = [], []\n    valid_loss = np.zeros(4,np.float32)\n    \n    for ep in range(n_epochs):\n        running_loss = 0\n        net = net.train()\n        print(f'Training for Epoch:{ep+1}')\n        for t, (images, targets) in tqdm(enumerate(train_loader), total=len(train_loader)):\n            batch_size = CFG.batch_size\n            images = images.cuda()\n            targets = targets.cuda()\n\n            with amp.autocast(enabled=True):\n                output = net(images)\n                loss0 = CFG.loss(output, targets.float())\n                \n            optimizer.zero_grad()\n            scaler.scale(loss0).backward()\n            scaler.unscale_(optimizer)\n            scaler.step(optimizer)\n            scheduler.step()\n            scaler.update()\n            \n            running_loss += loss0.item()\n            if t % 100 == 0:\n                print(f'Batch:{t} ; Train Batch Loss:{loss0.item():.4f}')\n        train_loss.append(running_loss / len(train_loader))\n        print(f'Epoch:{ep + 1}; Train Loss:{np.mean(train_loss):.4f}')\n        \n        valid_num = 0\n        valid_loss = 0\n        valid_probability = []\n        valid_truth = []\n\n        net = net.eval()\n        for t, (images, targets) in tqdm(enumerate(val_loader), total=len(val_loader)):\n            with torch.no_grad():\n                with amp.autocast(enabled=True):\n                    batch_size = CFG.batch_size\n                    images = images.cuda()\n                    targets = targets.cuda()\n                    output = net(images)\n                    loss0  = CFG.loss(output, targets.float())\n            \n            valid_num += batch_size\n            valid_loss += batch_size*loss0.item()\n            valid_truth.append(targets.data.cpu().numpy())\n            valid_probability.append(torch.sigmoid(output).data.cpu().numpy())\n            \n        truth = np.concatenate(valid_truth)\n        probability = np.concatenate(valid_probability)\n\n        loss = valid_loss/valid_num\n        metric = probabilistic_f1(truth, probability, beta=0.5)\n        auc = metrics.roc_auc_score(truth, probability)\n        \n        print(f'Epoch:{ep + 1}; Validation Loss:{loss:.4f}; F1 Score:{metric:.4f}; AUC Score:{auc:.4f}')\n        torch.save(net.state_dict(), f\"/kaggle/working/effnet_{ep + 1}.pth\")\n        torch.cuda.empty_cache()\n    return loss, metric, auc","metadata":{"papermill":{"duration":0.065041,"end_time":"2022-12-29T10:37:56.779551","exception":false,"start_time":"2022-12-29T10:37:56.71451","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:44:08.284256Z","iopub.execute_input":"2022-12-31T00:44:08.284641Z","iopub.status.idle":"2022-12-31T00:44:08.304822Z","shell.execute_reply.started":"2022-12-31T00:44:08.284607Z","shell.execute_reply":"2022-12-31T00:44:08.303553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss, f1, auc = train_loop(train_loader, val_loader)","metadata":{"papermill":{"duration":38402.78582,"end_time":"2022-12-29T21:17:59.578652","exception":false,"start_time":"2022-12-29T10:37:56.792832","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-12-31T00:44:09.295987Z","iopub.execute_input":"2022-12-31T00:44:09.296373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNATest(torch.utils.data.Dataset):\n    def __init__(self, test_df, transform):\n        self.df = pd.read_csv(test_df)\n        self.transform = transform\n        self.df['dcm_path'] = CFG.test_set + '/' + self.df.patient_id.apply(lambda x: str(x)) + '/' + self.df.image_id.apply(lambda x: str(x)) + '.dcm'\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path = self.df.dcm_path[index]\n        img = read_dicom(img_path)\n        img = img2roi(img, is_dicom=True)\n        img = cv2.resize(img, (512,512))\n        if self.transform:\n            self.transform(img)\n        img = torch.tensor(img, dtype=torch.float)\n        img = img.unsqueeze(0)\n        return img","metadata":{"execution":{"iopub.execute_input":"2022-12-29T21:18:02.829267Z","iopub.status.busy":"2022-12-29T21:18:02.828851Z","iopub.status.idle":"2022-12-29T21:18:02.837612Z","shell.execute_reply":"2022-12-29T21:18:02.836499Z"},"papermill":{"duration":1.646276,"end_time":"2022-12-29T21:18:02.840106","exception":false,"start_time":"2022-12-29T21:18:01.19383","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Infer(nn.Module):\n    def __init__(self, model):\n        super(Infer, self).__init__()\n        self.model = model\n    \n    def forward(self, batch):\n        x = self.model(batch)\n        cancer = torch.sigmoid(x)\n        cancer = torch.nan_to_num(cancer)\n        return cancer","metadata":{"execution":{"iopub.execute_input":"2022-12-29T21:18:17.763851Z","iopub.status.busy":"2022-12-29T21:18:17.76348Z","iopub.status.idle":"2022-12-29T21:18:17.769131Z","shell.execute_reply":"2022-12-29T21:18:17.76819Z"},"papermill":{"duration":1.413414,"end_time":"2022-12-29T21:18:17.771142","exception":false,"start_time":"2022-12-29T21:18:16.357728","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_model(test_loader):\n    model = Model(pretrained=False)\n    model.to(CFG.device)\n    model.load_state_dict(torch.load('/kaggle/input/eelan-weights/eelan_3.pth', map_location=CFG.device))\n    \n    model_infer = Infer(model)\n    model_infer.to(CFG.device)\n    preds = []\n    classes = []\n    with torch.no_grad():\n        model.eval()\n        for idx, batch in enumerate(test_loader):\n            batch = batch.to(CFG.device)\n            pred = model_infer(batch)\n            preds.append(pred.cpu().numpy())\n    return preds","metadata":{"execution":{"iopub.execute_input":"2022-12-29T21:18:20.777711Z","iopub.status.busy":"2022-12-29T21:18:20.777351Z","iopub.status.idle":"2022-12-29T21:18:20.786269Z","shell.execute_reply":"2022-12-29T21:18:20.783326Z"},"papermill":{"duration":1.464906,"end_time":"2022-12-29T21:18:20.788347","exception":false,"start_time":"2022-12-29T21:18:19.323441","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}