{"cells":[{"metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2021-01-31T14:33:42.832062Z","iopub.status.busy":"2021-01-31T14:33:42.8315Z","iopub.status.idle":"2021-01-31T14:33:44.652239Z","shell.execute_reply":"2021-01-31T14:33:44.652884Z"},"papermill":{"duration":1.838023,"end_time":"2021-01-31T14:33:44.653299","exception":false,"start_time":"2021-01-31T14:33:42.815276","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"import numpy as np \nimport pandas as pd\nimport os\nimport cv2\n\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as nnf","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:33:45.094639Z","iopub.status.busy":"2021-01-31T14:33:45.091416Z","iopub.status.idle":"2021-01-31T14:33:45.098016Z","shell.execute_reply":"2021-01-31T14:33:45.097478Z"},"papermill":{"duration":0.428918,"end_time":"2021-01-31T14:33:45.098207","exception":false,"start_time":"2021-01-31T14:33:44.669289","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"NUM_CL = 19\n\nBATCH = 16\nEPOCHS = 15\n\nLR = 0.0001\nIM_SIZE = 256\n\nDEVICE = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nPATH = '/kaggle/input/hpa-single-cell-image-classification/'\nTRAIN_DIR = PATH + 'train/'\nTEST_DIR = PATH + 'test/'","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:33:45.130148Z","iopub.status.busy":"2021-01-31T14:33:45.129599Z","iopub.status.idle":"2021-01-31T14:33:45.190306Z","shell.execute_reply":"2021-01-31T14:33:45.190718Z"},"papermill":{"duration":0.083612,"end_time":"2021-01-31T14:33:45.190862","exception":false,"start_time":"2021-01-31T14:33:45.10725","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"train = pd.read_csv(PATH +'train.csv')\n\n# I take just a subset to reduce training time \n# train = train[:1000]\n\ntrain.head()","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:33:45.215403Z","iopub.status.busy":"2021-01-31T14:33:45.213687Z","iopub.status.idle":"2021-01-31T14:33:45.216001Z","shell.execute_reply":"2021-01-31T14:33:45.216442Z"},"papermill":{"duration":0.016136,"end_time":"2021-01-31T14:33:45.216571","exception":false,"start_time":"2021-01-31T14:33:45.200435","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"Transform = transforms.Compose(\n    [transforms.ToTensor(),\n    transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))])","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:33:45.268714Z","iopub.status.busy":"2021-01-31T14:33:45.268143Z","iopub.status.idle":"2021-01-31T14:33:45.271024Z","shell.execute_reply":"2021-01-31T14:33:45.270626Z"},"papermill":{"duration":0.020332,"end_time":"2021-01-31T14:33:45.271141","exception":false,"start_time":"2021-01-31T14:33:45.250809","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"class GetData(Dataset):\n    def __init__(self, path, list_IDs, labels, img_size, Transform):\n        self.path = path\n        self.list_IDs = list_IDs\n        self.labels = labels\n        self.img_size = img_size        \n        self.transform = Transform\n        \n    def __len__(self):\n        return len(self.list_IDs)    \n    \n    def __getitem__(self, index):\n        ID = self.list_IDs[index]   \n        \n        # I take just a \"green\" images\n        data_file = cv2.imread(self.path + ID + '_green.png')\n            \n        img = cv2.resize(data_file, (self.img_size, self.img_size))\n        X = img/255.        \n        \n        if \"train\" in self.path:                       \n            y = self.labels[index]\n            y = y.split('|')\n            y = list(map(int, y))            \n            y = np.eye(NUM_CL, dtype='float')[y]                                    \n            y = y.sum(axis=0)            \n            return self.transform(X), y\n        \n        elif \"test\" in self.path:\n            return self.transform(X), ID","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:33:45.296468Z","iopub.status.busy":"2021-01-31T14:33:45.295218Z","iopub.status.idle":"2021-01-31T14:33:45.297597Z","shell.execute_reply":"2021-01-31T14:33:45.298007Z"},"papermill":{"duration":0.017976,"end_time":"2021-01-31T14:33:45.29814","exception":false,"start_time":"2021-01-31T14:33:45.280164","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"X_Train, Y_Train = train['ID'].values, train['Label'].values\n\ntrainset = GetData(TRAIN_DIR, X_Train, Y_Train, IM_SIZE, Transform)\ntrainloader = DataLoader(trainset, batch_size=BATCH, shuffle=True)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:33:45.320431Z","iopub.status.busy":"2021-01-31T14:33:45.319888Z","iopub.status.idle":"2021-01-31T14:33:45.383596Z","shell.execute_reply":"2021-01-31T14:33:45.383001Z"},"papermill":{"duration":0.076464,"end_time":"2021-01-31T14:33:45.383718","exception":false,"start_time":"2021-01-31T14:33:45.307254","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"X_Test = [name.rstrip('green.png').rstrip('_') for name in (os.listdir(TEST_DIR)) if '_green.png' in name]\n\ntestset = GetData(TEST_DIR, X_Test, None, IM_SIZE, Transform)\ntestloader = DataLoader(testset, batch_size=1, shuffle=False)","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:33:45.406991Z","iopub.status.busy":"2021-01-31T14:33:45.406379Z","iopub.status.idle":"2021-01-31T14:33:50.771387Z","shell.execute_reply":"2021-01-31T14:33:50.77083Z"},"papermill":{"duration":5.378172,"end_time":"2021-01-31T14:33:50.771526","exception":false,"start_time":"2021-01-31T14:33:45.393354","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"model = torchvision.models.resnet34()\nmodel.fc = nn.Linear(512, NUM_CL, bias=True)\nmodel = model.to(DEVICE)\n\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=LR)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# TRAIN"},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:33:50.800005Z","iopub.status.busy":"2021-01-31T14:33:50.799394Z","iopub.status.idle":"2021-01-31T14:34:13.127168Z","shell.execute_reply":"2021-01-31T14:34:13.126694Z"},"papermill":{"duration":22.3456,"end_time":"2021-01-31T14:34:13.127298","exception":false,"start_time":"2021-01-31T14:33:50.781698","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"%%time\n\nfor epoch in range(EPOCHS):\n    tr_loss = 0.0\n\n    model = model.train()\n\n    for i, (images, labels) in enumerate(trainloader):        \n        images = images.to(DEVICE)\n        labels = labels.to(DEVICE)       \n        logits = model(images.float())       \n        loss = criterion(logits, labels)\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        tr_loss += loss.detach().item()\n    \n    model.eval()\n    print('Epoch: %d | Loss: %.4f'%(epoch, tr_loss / i))","execution_count":null,"outputs":[]},{"metadata":{"papermill":{"duration":0.009985,"end_time":"2021-01-31T14:34:13.148322","exception":false,"start_time":"2021-01-31T14:34:13.138337","status":"completed"},"tags":[]},"cell_type":"markdown","source":"# TEST"},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:34:13.177948Z","iopub.status.busy":"2021-01-31T14:34:13.177379Z","iopub.status.idle":"2021-01-31T14:36:18.63088Z","shell.execute_reply":"2021-01-31T14:36:18.631333Z"},"papermill":{"duration":125.473148,"end_time":"2021-01-31T14:36:18.631496","exception":false,"start_time":"2021-01-31T14:34:13.158348","status":"completed"},"tags":[],"trusted":true},"cell_type":"code","source":"%%time\n\ns_ls = []\n\nwith torch.no_grad():\n    model.eval()\n    for image, fname in testloader:     \n        image = image.to(DEVICE)        \n        logits = model(image.float())                          \n        prob = nnf.softmax(logits, dim=1)\n        p, top_class = prob.topk(1, dim=1)\n        sp = ' '.join(str(e) for e in [top_class[0][0].item(), p[0][0].item()])               \n        img = cv2.imread(TEST_DIR + fname[0] + '_green.png')\n        \n        if img.shape[0] == 2048:\n            sp = sp + ' eNoLCAgIMAEABJkBdQ=='\n        elif img.shape[0] == 1728:\n            sp = sp + ' eNoLCAjJNgIABNkBkg=='\n        else:\n            sp = sp + ' eNoLCAgIsAQABJ4Beg=='\n        \n        s_ls.append([fname[0], img.shape[1], img.shape[0], sp])","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:36:18.664657Z","iopub.status.busy":"2021-01-31T14:36:18.664123Z","iopub.status.idle":"2021-01-31T14:36:18.670851Z","shell.execute_reply":"2021-01-31T14:36:18.670395Z"},"papermill":{"duration":0.02912,"end_time":"2021-01-31T14:36:18.670964","exception":false,"start_time":"2021-01-31T14:36:18.641844","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"sub = pd.DataFrame.from_records(s_ls, columns=['ID', 'ImageWidth', 'ImageHeight', 'PredictionString'])\n\nprint(len(sub))\nsub.head()","execution_count":null,"outputs":[]},{"metadata":{"execution":{"iopub.execute_input":"2021-01-31T14:36:18.696872Z","iopub.status.busy":"2021-01-31T14:36:18.696322Z","iopub.status.idle":"2021-01-31T14:36:18.840589Z","shell.execute_reply":"2021-01-31T14:36:18.839652Z"},"papermill":{"duration":0.158509,"end_time":"2021-01-31T14:36:18.840735","exception":false,"start_time":"2021-01-31T14:36:18.682226","status":"completed"},"tags":[],"trusted":false},"cell_type":"code","source":"sub.to_csv(\"submission.csv\", index=False)","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}