{"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 os, glob\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n!pip install ../input/timm-wheel/timm-0.6.5-py3-none-any.whl\n\nimport torch\nimport timm\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\nDATASET_FOLDER = \"/kaggle/input/mayo-clinic-strip-ai/\"\nDATASET_SMALL_FOLDER = \"/kaggle/input/stroke-blood-clot-origin-1k-scale-bg-crop\"","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:19:29.077255Z","iopub.execute_input":"2022-07-28T10:19:29.078412Z","iopub.status.idle":"2022-07-28T10:20:05.749114Z","shell.execute_reply.started":"2022-07-28T10:19:29.078308Z","shell.execute_reply":"2022-07-28T10:20:05.748103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path_csv = os.path.join(DATASET_FOLDER, \"train.csv\")\ndf_train = pd.read_csv(path_csv)\ndisplay(df_train.head())\n\nimport numpy as np\nfrom tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\n\ndef _color_means(img_path):\n    img = plt.imread(img_path)\n    if np.max(img) > 1.5:\n        img = img / 255.0\n    clr_mean = np.mean(img) if img.ndim == 2 else {i: np.mean(img[..., i]) for i in range(3)}\n    clr_std = np.std(img) if img.ndim == 2 else {i: np.std(img[..., i]) for i in range(3)}\n    return clr_mean, clr_std\n\n# os.path.join(DATASET_SMALL_FOLDER, \"train_images\")\nimages = glob.glob(os.path.join(DATASET_SMALL_FOLDER, \"train_images\", \"*.png\"))\nclr_mean_std = Parallel(n_jobs=os.cpu_count())(delayed(_color_means)(fn) for fn in tqdm(images[::10]))","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:20:05.751236Z","iopub.execute_input":"2022-07-28T10:20:05.752766Z","iopub.status.idle":"2022-07-28T10:20:10.789959Z","shell.execute_reply.started":"2022-07-28T10:20:05.752717Z","shell.execute_reply":"2022-07-28T10:20:10.788893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_color_mean = pd.DataFrame([c[0] for c in clr_mean_std]).describe()\ndisplay(img_color_mean.T)\nimg_color_std = pd.DataFrame([c[1] for c in clr_mean_std]).describe()\ndisplay(img_color_std.T)\n\nimg_color_mean = list(img_color_mean.T[\"mean\"])\nimg_color_std = list(img_color_std.T[\"mean\"])\nprint(img_color_mean, img_color_std)","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:20:10.791206Z","iopub.execute_input":"2022-07-28T10:20:10.791532Z","iopub.status.idle":"2022-07-28T10:20:10.851707Z","shell.execute_reply.started":"2022-07-28T10:20:10.791502Z","shell.execute_reply":"2022-07-28T10:20:10.850946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\n\nImage.MAX_IMAGE_PIXELS = 25_000_000_000\n\ndef prune_image_rows_cols(im, mask, thr=0.990):\n    # delete empty columns\n    for l in reversed(range(im.shape[1])):\n        if (np.sum(mask[:, l]) / float(mask.shape[0])) > thr:\n            im = np.delete(im, l, 1)\n    # delete empty rows\n    for l in reversed(range(im.shape[0])):\n        if (np.sum(mask[l, :]) / float(mask.shape[1])) > thr:\n            im = np.delete(im, l, 0)\n    return im\n\n\ndef mask_median(im, val=255):\n    masks = [None] * 3\n    for c in range(3):\n        masks[c] = im[..., c] >= np.median(im[:, :, c]) - 5\n    mask = np.logical_and(*masks)\n    im[mask, :] = val\n    return im, mask\n\n\ndef image_load_scale_norm(img_path, prune_thr=0.990, bg_val=255):\n    img = Image.open(img_path)\n    if (img.width * img.height) > 1_500_000_000:  # todo: for train images it was fine 4_000_000_000\n        print(img.width, img.height)\n        return None\n    scale = min(img.height / 2e3, img.width / 2e3)\n    tmp_size = int(img.width / scale), int(img.height / scale)\n    img.thumbnail(tmp_size, resample=Image.Resampling.BILINEAR, reducing_gap=scale)\n    im, mask = mask_median(np.array(img), val=bg_val)\n    im = prune_image_rows_cols(im, mask, thr=prune_thr)\n    img = Image.fromarray(im)\n    scale = min(img.height / 1e3, img.width / 1e3)\n    if scale > 1:\n        img = img.resize((int(img.width / scale), int(img.height / scale)), Image.LANCZOS)\n    return img","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:20:10.854402Z","iopub.execute_input":"2022-07-28T10:20:10.854712Z","iopub.status.idle":"2022-07-28T10:20:10.868363Z","shell.execute_reply.started":"2022-07-28T10:20:10.854684Z","shell.execute_reply":"2022-07-28T10:20:10.867101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nfrom tqdm.auto import tqdm\n\nls_imgs_tif = glob.glob(os.path.join(DATASET_FOLDER, \"test\", \"*.tif\"))\nnames = [os.path.splitext(os.path.basename(p))[0] for p in ls_imgs_tif]\npatient_ids = set([n.split(\"_\")[0] for n in names])\nprint(patient_ids)\n\n! mkdir test_images\n\nfor img_path in tqdm(ls_imgs_tif):\n    name, _ = os.path.splitext(os.path.basename(img_path))\n    img = image_load_scale_norm(img_path)\n    if not img:\n        print(f\"missing: {name}\")\n        continue\n    img.save(os.path.join(\"test_images\", f\"{name}.png\"))\n    del img\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:20:10.870477Z","iopub.execute_input":"2022-07-28T10:20:10.871329Z","iopub.status.idle":"2022-07-28T10:23:11.764725Z","shell.execute_reply.started":"2022-07-28T10:20:10.871284Z","shell.execute_reply":"2022-07-28T10:23:11.763404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train['img_name'] = df_train['image_id'].apply(lambda n: f\"{n}.png\")\ndisplay(df_train.head())\nprint(len(df_train))\nmissing = df_train['img_name'].apply(lambda n: not os.path.isfile(os.path.join(DATASET_SMALL_FOLDER, \"train_images\", n)))\ndf_train = df_train[~missing]\nprint(len(df_train))","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:23:11.766345Z","iopub.execute_input":"2022-07-28T10:23:11.766725Z","iopub.status.idle":"2022-07-28T10:23:12.119709Z","shell.execute_reply.started":"2022-07-28T10:23:11.766682Z","shell.execute_reply":"2022-07-28T10:23:12.117362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nfrom collections import Counter\nlabels = {\"CE\":0, \"LAA\":1}\n\nconverted_train_path = \"../input/stroke-blood-clot-origin-1k-scale-bg-crop/train_images/\"\n\ntrain_img_paths = []\ntrain_y = []\nfor n, (i,j) in enumerate(zip(df_train['image_id'].values, df_train['label'].values)):\n    print(f\"{n}/{len(df_train['label'].values)}\", end=\"\\r\")\n\n    train_img_paths.append(converted_train_path + i + '.png')\n\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 = 30\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:23:12.121623Z","iopub.execute_input":"2022-07-28T10:23:12.121978Z","iopub.status.idle":"2022-07-28T10:23:12.155437Z","shell.execute_reply.started":"2022-07-28T10:23:12.121945Z","shell.execute_reply":"2022-07-28T10:23:12.154413Z"},"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.ToTensor(),\n#            T.ConvertImageDtype(torch.float32),\n#            T.RandomCrop(size=(800, 800), pad_if_needed=True),\n#            T.GaussianBlur(kernel_size=5, sigma=(0.5, 4)),\n#            T.Resize((224, 224)),\n#            T.RandomHorizontalFlip(),\n#            T.RandomVerticalFlip(),\n#            T.RandomAffine(degrees=30, translate=(0.15,0.15)),\n#            T.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])])\n        \n        \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.Resize((224, 224)),\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.Resize((224, 224)),\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.open(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:23:12.157088Z","iopub.execute_input":"2022-07-28T10:23:12.157417Z","iopub.status.idle":"2022-07-28T10:23:12.173976Z","shell.execute_reply.started":"2022-07-28T10:23:12.157388Z","shell.execute_reply":"2022-07-28T10:23:12.173085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ViTBase16(nn.Module):\n    def __init__(self, n_classes, pretrained=False):\n\n        super(ViTBase16, self).__init__()\n\n        self.model = timm.create_model(\"vit_base_patch16_224\", pretrained=False)\n        if pretrained:\n            self.model.load_state_dict(torch.load(\"../input/vit-base-models-pretrained-pytorch/jx_vit_base_p16_224-80ecf9dd.pth\"))\n\n        self.model.head = nn.Linear(self.model.head.in_features, n_classes)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x\n\nepoch = 30\nvit_model = ViTBase16(2, pretrained=True)\nloss_history = [[], []] #train, val\naccuracy_history = [[], []] #train, val\nacc_epoch_history = [[],[]]\nloss_epoch_history = [[],[]]\n\noptimizer = torch.optim.Adam(vit_model.parameters(), lr=2e-04)\ncriterion = nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:24:15.011916Z","iopub.execute_input":"2022-07-28T10:24:15.013389Z","iopub.status.idle":"2022-07-28T10:24:21.994287Z","shell.execute_reply.started":"2022-07-28T10:24:15.013333Z","shell.execute_reply":"2022-07-28T10:24:21.992941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for e in range(epoch):\n    vit_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 = vit_model(data)\n        loss = criterion(output, target.view(-1,))\n        loss.backward()\n        \n        nn.utils.clip_grad_norm_(vit_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    vit_model.eval()\n    \n    with torch.no_grad():\n        for i, (data, target) in enumerate(validation_dataloader):\n            output = vit_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': vit_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:24:23.930696Z","iopub.execute_input":"2022-07-28T10:24:23.931167Z","iopub.status.idle":"2022-07-28T10:37:26.225008Z","shell.execute_reply.started":"2022-07-28T10:24:23.931128Z","shell.execute_reply":"2022-07-28T10:37:26.221973Z"},"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:37:26.227017Z","iopub.status.idle":"2022-07-28T10:37:26.227497Z","shell.execute_reply.started":"2022-07-28T10:37:26.22729Z","shell.execute_reply":"2022-07-28T10:37:26.227315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.save(vit_model.state_dict(), './vit16.pt')","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:37:26.230756Z","iopub.status.idle":"2022-07-28T10:37:26.231274Z","shell.execute_reply.started":"2022-07-28T10:37:26.230996Z","shell.execute_reply":"2022-07-28T10:37:26.23103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n\ntest_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(vit_model(transform_val(img).unsqueeze(0)))\n    del img\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:37:26.232648Z","iopub.status.idle":"2022-07-28T10:37:26.233463Z","shell.execute_reply.started":"2022-07-28T10:37:26.233237Z","shell.execute_reply":"2022-07-28T10:37:26.233261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prob_preds = nn.functional.softmax(torch.FloatTensor([i.detach().numpy() for i in pred]).view(-1,2), dim=1).numpy()\nprob_preds","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:37:26.235105Z","iopub.status.idle":"2022-07-28T10:37:26.235916Z","shell.execute_reply.started":"2022-07-28T10:37:26.235692Z","shell.execute_reply":"2022-07-28T10:37:26.235716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"prob = pd.DataFrame({\"CE\" : prob_preds[:,0], \"LAA\" : prob_preds[:,1], \"id\" : test_names}).groupby(\"id\").mean()\n\nsubmission = pd.read_csv(\"../input/mayo-clinic-strip-ai/sample_submission.csv\")\n\nsubmission.CE = prob.CE.to_list()\nsubmission.LAA = prob.LAA.to_list()\nsubmission","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:37:26.237422Z","iopub.status.idle":"2022-07-28T10:37:26.237841Z","shell.execute_reply.started":"2022-07-28T10:37:26.237635Z","shell.execute_reply":"2022-07-28T10:37:26.237655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m = torch.jit.script(vit_model)\n\n# Save to file\ntorch.jit.save(m, 'vit_model16_2e4.pt')","metadata":{},"execution_count":null,"outputs":[]}]}