{"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\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\nif \"converted\" not in os.listdir():\n    os.mkdir(\"converted/\")\n\nfrom 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:56:21.032849Z","iopub.execute_input":"2022-07-28T10:56:21.033255Z","iopub.status.idle":"2022-07-28T10:56:52.724736Z","shell.execute_reply.started":"2022-07-28T10:56:21.033225Z","shell.execute_reply":"2022-07-28T10:56:52.723171Z"},"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:56:52.727716Z","iopub.execute_input":"2022-07-28T10:56:52.728216Z","iopub.status.idle":"2022-07-28T10:56:52.758593Z","shell.execute_reply.started":"2022-07-28T10:56:52.728166Z","shell.execute_reply":"2022-07-28T10:56:52.7577Z"},"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:56:52.759774Z","iopub.execute_input":"2022-07-28T10:56:52.760588Z","iopub.status.idle":"2022-07-28T10:56:55.895276Z","shell.execute_reply.started":"2022-07-28T10:56:52.760556Z","shell.execute_reply":"2022-07-28T10:56:55.894398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = {\"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        np.save(f\"./converted/{i}.npy\", cv2.resize(tifi.imread(train_path + i + '.tif'),(224, 224)))\n        train_img_paths.append(converted_train_path + i + '.npy')\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:56:55.897342Z","iopub.execute_input":"2022-07-28T10:56:55.897733Z","iopub.status.idle":"2022-07-28T10:56:56.729008Z","shell.execute_reply.started":"2022-07-28T10:56:55.897692Z","shell.execute_reply":"2022-07-28T10:56:56.727805Z"},"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:56:56.730697Z","iopub.execute_input":"2022-07-28T10:56:56.731028Z","iopub.status.idle":"2022-07-28T10:56:56.749083Z","shell.execute_reply.started":"2022-07-28T10:56:56.730991Z","shell.execute_reply":"2022-07-28T10:56:56.748134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class xcit24(nn.Module):\n    def __init__(self, n_classes, pretrained=False):\n\n        super(xcit24, self).__init__()\n\n        self.model = timm.create_model(\"xcit_medium_24_p16_224\", pretrained=False)\n        if pretrained:\n            self.model.load_state_dict(torch.load(\"../input/timm-convnext-xcit/xcit_medium_24_p16_224.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 = 20\nxcit_model = xcit24(2, True)\nloss_history = [[], []] #train, val\naccuracy_history = [[], []] #train, val\nacc_epoch_history = [[],[]]\nloss_epoch_history = [[],[]]\n\noptimizer = torch.optim.Adam(xcit_model.parameters(), lr=2e-04)\ncriterion = nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:57:15.683123Z","iopub.execute_input":"2022-07-28T10:57:15.683535Z","iopub.status.idle":"2022-07-28T10:57:20.764743Z","shell.execute_reply.started":"2022-07-28T10:57:15.683502Z","shell.execute_reply":"2022-07-28T10:57:20.763611Z"},"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:56:56.807832Z","iopub.status.idle":"2022-07-28T10:56:56.808231Z","shell.execute_reply.started":"2022-07-28T10:56:56.808045Z","shell.execute_reply":"2022-07-28T10:56:56.808063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for e in range(epoch):\n    xcit_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 = xcit_model(data)\n        loss = criterion(output, target.view(-1,))\n        loss.backward()\n        \n        nn.utils.clip_grad_norm_(xcit_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    xcit_model.eval()\n    \n    with torch.no_grad():\n        for i, (data, target) in enumerate(validation_dataloader):\n            output = xcit_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': xcit_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:23.134149Z","iopub.execute_input":"2022-07-28T10:57:23.134533Z","iopub.status.idle":"2022-07-28T11:04:01.30769Z","shell.execute_reply.started":"2022-07-28T10:57:23.134502Z","shell.execute_reply":"2022-07-28T11:04:01.305206Z"},"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:56:56.811951Z","iopub.status.idle":"2022-07-28T10:56:56.813356Z","shell.execute_reply.started":"2022-07-28T10:56:56.812894Z","shell.execute_reply":"2022-07-28T10:56:56.812927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_path = '../input/mayo-clinic-strip-ai/test/'\npred = []\ntransform_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\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(xcit_model(transform_val(img).unsqueeze(0)))","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:56:56.817429Z","iopub.status.idle":"2022-07-28T10:56:56.818205Z","shell.execute_reply.started":"2022-07-28T10:56:56.817894Z","shell.execute_reply":"2022-07-28T10:56:56.817925Z"},"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:56:56.820298Z","iopub.status.idle":"2022-07-28T10:56:56.821581Z","shell.execute_reply.started":"2022-07-28T10:56:56.821234Z","shell.execute_reply":"2022-07-28T10:56:56.821267Z"},"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:56:56.823914Z","iopub.status.idle":"2022-07-28T10:56:56.824545Z","shell.execute_reply.started":"2022-07-28T10:56:56.824227Z","shell.execute_reply":"2022-07-28T10:56:56.824256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_names","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:56:56.825603Z","iopub.status.idle":"2022-07-28T10:56:56.825974Z","shell.execute_reply.started":"2022-07-28T10:56:56.825796Z","shell.execute_reply":"2022-07-28T10:56:56.825813Z"},"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:56:56.827136Z","iopub.status.idle":"2022-07-28T10:56:56.827545Z","shell.execute_reply.started":"2022-07-28T10:56:56.827329Z","shell.execute_reply":"2022-07-28T10:56:56.827346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:56:56.828832Z","iopub.status.idle":"2022-07-28T10:56:56.829213Z","shell.execute_reply.started":"2022-07-28T10:56:56.829031Z","shell.execute_reply":"2022-07-28T10:56:56.829049Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m = torch.jit.script(xcit_model)\n\n# Save to file\ntorch.jit.save(m, 'xcit_model.pt')\n","metadata":{"execution":{"iopub.status.busy":"2022-07-28T10:56:56.830279Z","iopub.status.idle":"2022-07-28T10:56:56.831975Z","shell.execute_reply.started":"2022-07-28T10:56:56.831778Z","shell.execute_reply":"2022-07-28T10:56:56.8318Z"},"trusted":true},"execution_count":null,"outputs":[]}]}