{"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 ../input/timm-wheel/timm-0.6.5-py3-none-any.whl\n!conda install ../input/how-to-use-pyvips-offline/*.tar.bz2 \n\nimport torch\nimport timm\nimport pyvips\nimport pandas as pd\nimport os\nimport tifffile as tifi\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport random\nimport cv2\nimport h5py\nimport pandas as pd\nimport torchvision.models as models\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport torch.nn as nn\nimport torch.nn.functional as F\n\ntorch.manual_seed(0)\nrandom.seed(0)\nnp.random.seed(0)\n\nif \"converted\" not in os.listdir():\n    os.mkdir(\"converted/\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-28T10:56:01.813832Z","iopub.execute_input":"2022-07-28T10:56:01.81435Z","iopub.status.idle":"2022-07-28T10:57:22.197412Z","shell.execute_reply.started":"2022-07-28T10:56:01.814317Z","shell.execute_reply":"2022-07-28T10:57:22.196136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def vips2numpy(vi):\n    \n    # map vips formats to np dtypes\n    format_to_dtype = {\n        'uchar': np.uint8,       'char': np.int8,\n        'ushort': np.uint16,     'short': np.int16,\n        'uint': np.uint32,       'int': np.int32,\n        'float': np.float32,     'double': np.float64,\n        'complex': np.complex64, 'dpcomplex': np.complex128,\n    }\n    \n    # Return newly written np.ndarray\n    return np.ndarray(buffer=vi.write_to_memory(),\n                      dtype=format_to_dtype[vi.format],\n                      shape=[vi.height, vi.width, vi.bands])\n\ndef pyvips_open_downsampled_slide(img_path, downsample_by=8, as_numpy=True, resize_to=(512,512)):\n    \"\"\"\n    \n    Helper function to convert WSI into smaller downscaled version using pyvips.\n    \n    Timing details for MAYO CLINIC STRIP AI dataset:\n        SMALLEST IMAGE BY AREA (4417, 5314)\n            * Function takes ~1 seconds to run\n        MEDIAN IMAGE BY AREA   (17573, 38743)(~30X LARGER THAN SMALLEST IMAGE)\n            * Function takes ~30 seconds to run\n        LARGEST IMAGE BY AREA  (48282, 101406)(~208X LARGER THAN SMALLEST IMAGE)(~7X LARGER THAN MEDIAN IMAGE)\n            * Function takes ~405 seconds to run\n    \n    Args:\n        img_path (str): Path to .tif file to be downsampled\n        downsample_by (int): How many times smaller should resultant \n            image be. i.e. image_size*(1/downsample_by) = new_size\n            -1 will yield maximum downsample above resized image shape\n        as_numpy (bool, optional): Whether to return image as numpy array (default)\n           or leave as PIL.Image object for further manipulation\n        resize_to (tuple of ints, optional): What to resize the downsampled image to\n    \n    Returns:\n        Downsampled image as a numpy array of type uint8 with only 3 channels\n    \"\"\"\n    \n    # Open the image with PIL\n    tmp_img = pyvips.Image.new_from_file(img_path)    \n    \n    print(\"\\n... APPROXIMATE TIME TO LOAD IMAGE IS AT MOST APPROXIMATELY: \" \\\n          f\"{int((405/(48282*101406))*(tmp_img.width*tmp_img.height))} SECONDS ...\\n\")\n    \n    # if -1 than we downsample by whatever results in the image having dimensions as\n    # close to 512x512 as possible so the image can be resized after\n    \n    if downsample_by==-1:\n        _epsilon = 1e-3\n        downsample_by=min(tmp_img.width, tmp_img.height)/resize_to[0]-_epsilon\n    \n    # Resize the image\n    tmp_img = tmp_img.resize(1/downsample_by)\n    tmp_img = vips2numpy(tmp_img) if as_numpy else tmp_img\n    tmp_img = cv2.resize(tmp_img, resize_to) if resize_to is not None else tmp_img\n    \n    return tmp_img","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:57:22.200002Z","iopub.execute_input":"2022-07-28T10:57:22.200452Z","iopub.status.idle":"2022-07-28T10:57:22.217091Z","shell.execute_reply.started":"2022-07-28T10:57:22.200413Z","shell.execute_reply":"2022-07-28T10:57:22.215862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections  import Counter\nCounter(pd.read_csv(\"../input/mayo-clinic-strip-ai/train.csv\").label.values.tolist())","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:57:22.218575Z","iopub.execute_input":"2022-07-28T10:57:22.219591Z","iopub.status.idle":"2022-07-28T10:57:22.335725Z","shell.execute_reply.started":"2022-07-28T10:57:22.219542Z","shell.execute_reply":"2022-07-28T10:57:22.334852Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"../input/mayo-clinic-strip-ai/train.csv\")\ntrain_df              ","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:57:22.338065Z","iopub.execute_input":"2022-07-28T10:57:22.33863Z","iopub.status.idle":"2022-07-28T10:57:22.366167Z","shell.execute_reply.started":"2022-07-28T10:57:22.338597Z","shell.execute_reply":"2022-07-28T10:57:22.365305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\ntrain_path = \"../input/mayo-clinic-strip-ai/train/\"\n\nimg = Image.open(\"../input/mayo-clinic-strip-ai/train/026c97_0.tif\")\n\nx1 = img.size[0]\nx2 = img.size[1]\nsc = x2/x1\n\n\n\nfactor = 1/10\nresized_img = img.resize((int(x1*factor), int(sc*x1*factor)))\n\nprint(resized_img.size)\nresized_img.rotate(90, expand=True)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:57:22.3675Z","iopub.execute_input":"2022-07-28T10:57:22.368045Z","iopub.status.idle":"2022-07-28T10:57:26.446438Z","shell.execute_reply.started":"2022-07-28T10:57:22.368012Z","shell.execute_reply":"2022-07-28T10:57:26.444972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\ngc.enable()\n\nlabels = {\"CE\":0, \"LAA\":1}\n\nconverted_train_path = \"./converted/\"\nsaved_converted_path = \"../input/mayo-clinic-strip-ainumpy-files/\"\ntrain_img_paths = []\ntrain_y = []\nfor n, (i,j) in enumerate(zip(train_df['image_id'].values, train_df['label'].values)):\n    print(f\"{n}/{len(train_df['label'].values)}\", end=\"\\r\")\n    if f\"{i}.npy\" not in os.listdir(saved_converted_path):\n        img = pyvips_open_downsampled_slide(test_path+i+'.tif', downsample_by=-1, as_numpy=True, resize_to=(384,384))\n        np.save(f\"./converted/{i}.npy\", img)\n        train_img_paths.append(converted_train_path + i + '.npy')\n        del img\n        gc.collect()\n                \n    else:\n        train_img_paths.append(saved_converted_path + i + '.npy')\n                \n    train_y.append(labels[j])\n\nc = list(zip(train_img_paths, train_y))\n\nrandom.shuffle(c)\n\ntrain_img_paths, train_y = zip(*c)\ntrain_y = list(train_y)\ntrain_img_paths = list(train_img_paths)\n\nn_val = 40\n\nval_y = []\nval_img_paths = []\n\ni = 0 # start\nj = 0 # class 0 count\nk = 0 # class 1 count\n\n\nwhile len(val_y)!=n_val:\n    if train_y[i] == 0 and j<int(n_val/2):\n        j+=1\n        val_y.append(train_y.pop(i))\n        val_img_paths.append(train_img_paths.pop(0))\n        i=0\n        \n\n    elif train_y[i] == 1 and k<n_val-int(n_val/2):\n        k+=1\n        val_y.append(train_y.pop(i))\n        val_img_paths.append(train_img_paths.pop(0))\n        i=0\n        \n    else:\n        i+=1\n        \nCounter(val_y)","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:57:26.448048Z","iopub.execute_input":"2022-07-28T10:57:26.448432Z","iopub.status.idle":"2022-07-28T10:57:27.27816Z","shell.execute_reply.started":"2022-07-28T10:57:26.448399Z","shell.execute_reply":"2022-07-28T10:57:27.277137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImgDataset(Dataset):\n    def __init__(self, x, y, dataset_type):\n        self.x = x\n        self.y = y\n        self.transform_train = T.Compose([T.RandomHorizontalFlip(),\n                                    T.RandomVerticalFlip(),\n                                    T.RandomRotation(30),\n                                    T.ToTensor(),\n                                    T.ConvertImageDtype(torch.float32),\n                                    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])\n   \n        self.transform_val = T.Compose([T.ToTensor(),\n                                    T.ConvertImageDtype(torch.float32),\n                                    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])\n    \n        self.dataset_type = dataset_type\n        \n#         if 'cache' not in os.listdir('./'):\n#             h5_file = h5py.File('cache', 'w')\n        \n#         else:\n#             h5_file = h5py.File('cache', 'a')\n        \n# #         if self.dataset_type not in h5_file:\n# #             index_dataset = h5_file.create_dataset(dataset_type, shape=(len(x), 3, 224, 224), dtype=np.float32, fillvalue=0)\n#         h5_file.close()\n        \n    def __len__(self):\n        return len(self.x)\n    \n    def __getitem__(self, idx):\n#         h5_file = h5py.File('cache', 'a')\n        \n#         if f'{idx}_{self.dataset_type}' not in h5_file: # true if empty/not cached\n#             img = Image.fromarray(cv2.resize(tifi.imread(self.x[idx]),(224, 224)))\n#             if self.dataset_type == 'train':\n#                 img = self.transform_train(img)\n#             elif self.dataset_type == 'val':\n#                 img = self.transform_val(img)\n                \n#             index_dataset = h5_file.create_dataset(f'{idx}_{self.dataset_type}', shape=(3, 224, 224), dtype=np.float32, data = img.numpy(), chunks=True)\n            \n#         else:\n#             img = torch.FloatTensor(data[f'{idx}_{self.dataset_type}'])\n\n        img = Image.fromarray(np.load(self.x[idx]))\n        if self.dataset_type == 'train':\n            img = self.transform_train(img)\n        elif self.dataset_type == 'val':\n            img = self.transform_val(img)\n            \n        label = torch.LongTensor([self.y[idx]])\n#         h5_file.close()\n        \n\n        return img, label \n    \n    \ndef seed_worker(worker_id):\n    worker_seed = torch.initial_seed() % 2**32\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)\n\ng = torch.Generator()\ng.manual_seed(0)\n    \ntrain_dataset = ImgDataset(train_img_paths, train_y, dataset_type = 'train')\nvalidation_dataset = ImgDataset(val_img_paths, val_y, dataset_type = 'val')\nvalidation_dataloader = DataLoader(validation_dataset, batch_size = 16, shuffle=False, worker_init_fn=seed_worker, generator=g)\ntrain_dataloader = DataLoader(train_dataset, batch_size = 16, shuffle=False, worker_init_fn=seed_worker, generator=g)","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:57:27.279617Z","iopub.execute_input":"2022-07-28T10:57:27.279938Z","iopub.status.idle":"2022-07-28T10:57:27.296934Z","shell.execute_reply.started":"2022-07-28T10:57:27.27991Z","shell.execute_reply":"2022-07-28T10:57:27.29592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ConvNext(nn.Module):\n    def __init__(self, n_classes, pretrained=True):\n\n        super(ConvNext, self).__init__()\n\n        self.model = timm.create_model(\"convnext_small_384_in22ft1k\", pretrained=False)\n        if pretrained:\n            self.model.load_state_dict(torch.load(\"../input/timm-convnext-xcit/convnext_small_384_in22ft1k.pth\"))\n        self.model.head.fc = nn.Linear(self.model.head.fc.in_features, n_classes)\n        \n    def forward(self, x):\n        x = self.model(x)\n        return x\n\nepoch = 20\nconvnext_model = ConvNext(2, True)\nloss_history = [[], []] #train, val\naccuracy_history = [[], []] #train, val\nacc_epoch_history = [[],[]]\nloss_epoch_history = [[],[]]\n\noptimizer = torch.optim.Adam(convnext_model.parameters(), lr=2e-04)\ncriterion = nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T11:03:08.643752Z","iopub.execute_input":"2022-07-28T11:03:08.644192Z","iopub.status.idle":"2022-07-28T11:03:09.697062Z","shell.execute_reply.started":"2022-07-28T11:03:08.644156Z","shell.execute_reply":"2022-07-28T11:03:09.696033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Caching\n# for i, (data, target) in enumerate(train_dataloader):\n#     print(f\"MINIBATCH {i+1}/{train_dataloader.__len__()}\")\n    \n# for i, (data, target) in enumerate(val_dataloader):\n#     print(f\"MINIBATCH {i+1}/{val_dataloader.__len__()}\")","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:57:28.690602Z","iopub.execute_input":"2022-07-28T10:57:28.69107Z","iopub.status.idle":"2022-07-28T10:57:28.696745Z","shell.execute_reply.started":"2022-07-28T10:57:28.691024Z","shell.execute_reply":"2022-07-28T10:57:28.695616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for e in range(epoch):\n    convnext_model.train()\n    print(f\"====================== EPOCH {e+1} ======================\")\n    print(\"Training.....\")\n    for i, (data, target) in enumerate(train_dataloader):\n        optimizer.zero_grad()\n        output = convnext_model(data)\n        loss = criterion(output, target.view(-1,))\n        loss.backward()\n        \n        nn.utils.clip_grad_norm_(convnext_model.parameters(), 3)\n        accuracy = (output.argmax(dim=1) == target).float().mean()\n        \n        loss_history[0].append(loss.item())\n        accuracy_history[0].append(accuracy)\n        \n        optimizer.step()\n        \n        print(f\"MINIBATCH {i+1}/{train_dataloader.__len__()} TRAIN ACC : {accuracy_history[0][-1]}  TRAIN LOSS : {loss_history[0][-1]}\")\n            \n    \n    print(\"Validation.....\")\n    convnext_model.eval()\n    \n    with torch.no_grad():\n        for i, (data, target) in enumerate(validation_dataloader):\n            output = convnext_model(data)\n            loss = criterion(output, target.view(-1,))\n            accuracy = (output.argmax(dim=1) == target).float().mean()\n            loss_history[1].append(loss.item())\n            accuracy_history[1].append(accuracy)\n        \n    acc_epoch_history[0].append(sum(accuracy_history[0][-1:-train_dataloader.__len__():-1])/train_dataloader.__len__())\n    acc_epoch_history[1].append(sum(accuracy_history[1][-1:-validation_dataloader.__len__():-1])/validation_dataloader.__len__())\n    \n    loss_epoch_history[0].append(sum(loss_history[0][-1:-train_dataloader.__len__():-1])/train_dataloader.__len__())\n    loss_epoch_history[1].append(sum(loss_history[1][-1:-validation_dataloader.__len__():-1])/validation_dataloader.__len__())\n    \n    print(\"====================================================\")\n    print(f\"TRAIN ACC : {acc_epoch_history[0][-1]}  TRAIN LOSS : {loss_epoch_history[0][-1]}\")\n    print(f\"VALL ACC : {acc_epoch_history[1][-1]}  VAL LOSS : {loss_epoch_history[1][-1]}\")\n    print(\"====================================================\")\n    \n    torch.save({\n            'epoch': e,\n            'model_state_dict': convnext_model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'loss': loss_epoch_history[0][-1],\n            'acc' : acc_epoch_history[0][-1]\n            }, './model_checkpoint.pt')","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:57:28.701617Z","iopub.execute_input":"2022-07-28T10:57:28.702354Z","iopub.status.idle":"2022-07-28T10:58:35.387474Z","shell.execute_reply.started":"2022-07-28T10:57:28.702306Z","shell.execute_reply":"2022-07-28T10:58:35.385813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplt.plot(acc_epoch_history[0], label=\"train\")\nplt.plot(acc_epoch_history[1], label=\"val\")\nplt.legend()\nplt.title('ACCURACY VS EPOCH')\nplt.show()\n\nplt.plot(loss_epoch_history[0], label=\"train\")\nplt.plot(loss_epoch_history[1], label=\"val\")\nplt.legend()\nplt.title('LOSS VS EPOCH')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:58:35.389033Z","iopub.status.idle":"2022-07-28T10:58:35.389522Z","shell.execute_reply.started":"2022-07-28T10:58:35.389309Z","shell.execute_reply":"2022-07-28T10:58:35.38933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_path = '../input/mayo-clinic-strip-ai/test/'\npred = []\ntransform_val = T.Compose([T.PILToTensor(),\n                                    T.ConvertImageDtype(torch.float32),\n                                    T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])\n\ntest_names = list(os.listdir(test_path))\nfor i in test_names:\n    img = Image.fromarray(cv2.resize(tifi.imread(test_path+i),(224, 224)))\n    pred.append(convnext_model(transform_val(img).unsqueeze(0)))","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:58:35.394371Z","iopub.status.idle":"2022-07-28T10:58:35.394981Z","shell.execute_reply.started":"2022-07-28T10:58:35.394694Z","shell.execute_reply":"2022-07-28T10:58:35.394722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = nn.functional.softmax(torch.FloatTensor([i.detach().numpy() for i in pred]).view(-1,2), dim=1).numpy()\nsubmission","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:58:35.39703Z","iopub.status.idle":"2022-07-28T10:58:35.397633Z","shell.execute_reply.started":"2022-07-28T10:58:35.397335Z","shell.execute_reply":"2022-07-28T10:58:35.397363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_csv = pd.read_csv('../input/mayo-clinic-strip-ai/sample_submission.csv')\nsubmission_csv","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:58:35.399332Z","iopub.status.idle":"2022-07-28T10:58:35.399935Z","shell.execute_reply.started":"2022-07-28T10:58:35.399643Z","shell.execute_reply":"2022-07-28T10:58:35.399671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_names","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:58:35.401853Z","iopub.status.idle":"2022-07-28T10:58:35.402486Z","shell.execute_reply.started":"2022-07-28T10:58:35.402163Z","shell.execute_reply":"2022-07-28T10:58:35.402192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = {'patient_id':[i[:-6] for i in test_names],\n        'CE':submission[:,0].tolist(),\n        'LAA':submission[:,1].tolist()}\n  \ndf = pd.DataFrame(data)\n\ndf","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:58:35.405143Z","iopub.status.idle":"2022-07-28T10:58:35.405628Z","shell.execute_reply.started":"2022-07-28T10:58:35.405379Z","shell.execute_reply":"2022-07-28T10:58:35.405398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:58:35.408248Z","iopub.status.idle":"2022-07-28T10:58:35.410356Z","shell.execute_reply.started":"2022-07-28T10:58:35.410094Z","shell.execute_reply":"2022-07-28T10:58:35.410139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m = torch.jit.script(convnext_model)\n\n# Save to file\ntorch.jit.save(m, 'convnext_model.pt')\n","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:58:35.411754Z","iopub.status.idle":"2022-07-28T10:58:35.412626Z","shell.execute_reply.started":"2022-07-28T10:58:35.412374Z","shell.execute_reply":"2022-07-28T10:58:35.412401Z"},"trusted":true},"execution_count":null,"outputs":[]}]}