{"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":"markdown","source":"# **Imports**","metadata":{"papermill":{"duration":0.012363,"end_time":"2021-10-26T02:08:38.871395","exception":false,"start_time":"2021-10-26T02:08:38.859032","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import os\nimport cv2\nimport warnings\nimport random\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nfrom torch import nn\nfrom torch import optim\nfrom torch.nn import CrossEntropyLoss\nfrom torch.utils.data import DataLoader, Dataset\n\nfrom albumentations import Normalize, Resize, Compose\nfrom albumentations.pytorch import ToTensorV2\n\nwarnings.filterwarnings(\"ignore\")\n\ndef fix_all_seeds(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nfix_all_seeds(42)","metadata":{"papermill":{"duration":3.930151,"end_time":"2021-10-26T02:08:42.81341","exception":false,"start_time":"2021-10-26T02:08:38.883259","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-11-04T18:46:11.484738Z","iopub.execute_input":"2021-11-04T18:46:11.485005Z","iopub.status.idle":"2021-11-04T18:46:11.494279Z","shell.execute_reply.started":"2021-11-04T18:46:11.484977Z","shell.execute_reply":"2021-11-04T18:46:11.493457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SAMPLE_SUBMISSION  = '../input/sartorius-cell-instance-segmentation/sample_submission.csv'\nTRAIN_CSV = \"../input/sartorius-cell-instance-segmentation/train.csv\"\nTRAIN_PATH = \"../input/sartorius-cell-instance-segmentation/train\"\nTEST_PATH = \"../input/sartorius-cell-instance-segmentation/test\"\nEXTRA_DATA_PATH = \"../input/sartorius-cell-instance-segmentation/train_semi_supervised\"","metadata":{"papermill":{"duration":0.018807,"end_time":"2021-10-26T02:08:42.900106","exception":false,"start_time":"2021-10-26T02:08:42.881299","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-11-04T18:46:11.496153Z","iopub.execute_input":"2021-11-04T18:46:11.496400Z","iopub.status.idle":"2021-11-04T18:46:11.507692Z","shell.execute_reply.started":"2021-11-04T18:46:11.496368Z","shell.execute_reply":"2021-11-04T18:46:11.506980Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(TRAIN_CSV)\ndf_class = df_train.groupby(\"id\")[['cell_type']].first().reset_index()\ndf_class['cell_type'].value_counts(normalize=True).round(2)\n\ndf_class_train, df_class_val = train_test_split(df_class, test_size=0.2)\ndf_class_train['cell_type'].value_counts(normalize=True).round(2)","metadata":{"execution":{"iopub.status.busy":"2021-11-04T18:46:11.509173Z","iopub.execute_input":"2021-11-04T18:46:11.509666Z","iopub.status.idle":"2021-11-04T18:46:11.820357Z","shell.execute_reply.started":"2021-11-04T18:46:11.509631Z","shell.execute_reply":"2021-11-04T18:46:11.819516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Efficientnet Model**","metadata":{"papermill":{"duration":0.013163,"end_time":"2021-10-26T02:08:42.838836","exception":false,"start_time":"2021-10-26T02:08:42.825673","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!pip install efficientnet_pytorch","metadata":{"execution":{"iopub.status.busy":"2021-11-04T18:46:11.821873Z","iopub.execute_input":"2021-11-04T18:46:11.822137Z","iopub.status.idle":"2021-11-04T18:46:18.857443Z","shell.execute_reply.started":"2021-11-04T18:46:11.822102Z","shell.execute_reply":"2021-11-04T18:46:18.856607Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from efficientnet_pytorch import EfficientNet\nmodel = EfficientNet.from_pretrained(model_name='efficientnet-b3', num_classes=3)","metadata":{"execution":{"iopub.status.busy":"2021-11-04T18:46:18.860459Z","iopub.execute_input":"2021-11-04T18:46:18.860752Z","iopub.status.idle":"2021-11-04T18:46:19.081528Z","shell.execute_reply.started":"2021-11-04T18:46:18.860700Z","shell.execute_reply":"2021-11-04T18:46:19.080782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Semi-Supervised Data**","metadata":{"papermill":{"duration":0.013309,"end_time":"2021-10-26T02:08:43.639894","exception":false,"start_time":"2021-10-26T02:08:43.626585","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CellClassificationDatasetExtraData(Dataset):\n    def __init__(self):\n        self.base_path = EXTRA_DATA_PATH\n        self.transforms = Compose([\n            Resize(224, 224), \n            Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), p=1), \n            ToTensorV2()\n        ])\n        self.files = os.listdir(EXTRA_DATA_PATH)\n        self.labels = ['shsy5y', 'astro', 'cort']\n\n\n    def __getitem__(self, idx):\n        file = self.files[idx]\n        image_path = os.path.join(self.base_path, file)\n        image = self.transforms(image=cv2.imread(image_path))['image']\n        \n        label = file.split(\"[\")[0]\n        if label == 'astros':\n            label = 'astro'\n            \n        return {'image': image, 'label': self.labels.index(label)}\n\n    def __len__(self):\n        return len(self.files)","metadata":{"papermill":{"duration":0.023679,"end_time":"2021-10-26T02:08:43.676944","exception":false,"start_time":"2021-10-26T02:08:43.653265","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-11-04T18:46:19.082996Z","iopub.execute_input":"2021-11-04T18:46:19.083403Z","iopub.status.idle":"2021-11-04T18:46:19.092026Z","shell.execute_reply.started":"2021-11-04T18:46:19.083365Z","shell.execute_reply":"2021-11-04T18:46:19.091204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Dataset**","metadata":{"papermill":{"duration":0.013772,"end_time":"2021-10-26T02:08:43.703925","exception":false,"start_time":"2021-10-26T02:08:43.690153","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CellClassificationDataset(Dataset):\n    def __init__(self, df):\n        self.df = df\n        self.base_path = TRAIN_PATH\n        self.transforms = Compose([\n            Resize(244, 244), \n            Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), p=1), \n            ToTensorV2()\n        ])\n        self.image_ids = df.id.unique().tolist()\n        self.labels = ['shsy5y', 'astro', 'cort']\n\n    def get_label_for_img(self, image_id):\n        label = self.df.loc[self.df['id'] == image_id, 'cell_type'].iloc[0]\n        label_id = self.labels.index(label)\n        return label_id\n        \n    def __getitem__(self, idx):\n        image_id = self.image_ids[idx]\n        image_path = os.path.join(self.base_path, image_id + \".png\")\n        image = self.transforms(image=cv2.imread(image_path))['image']\n        label = self.get_label_for_img(image_id)\n        return {'image': image, 'label': label}\n\n    def __len__(self):\n        return len(self.image_ids)","metadata":{"papermill":{"duration":0.025154,"end_time":"2021-10-26T02:08:43.742952","exception":false,"start_time":"2021-10-26T02:08:43.717798","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-11-04T18:47:04.178230Z","iopub.execute_input":"2021-11-04T18:47:04.178568Z","iopub.status.idle":"2021-11-04T18:47:04.195268Z","shell.execute_reply.started":"2021-11-04T18:47:04.178527Z","shell.execute_reply":"2021-11-04T18:47:04.191864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_train = CellClassificationDataset(df_class_train)\ndl_train = DataLoader(ds_train, batch_size=64, num_workers=8, pin_memory=True, shuffle=True)\n\nds_train_extra = CellClassificationDatasetExtraData()\ndl_train_extra = DataLoader(ds_train_extra, batch_size=64, num_workers=8, pin_memory=True, shuffle=True)\n\nds_val = CellClassificationDataset(df_class_val)\ndl_val = DataLoader(ds_val, batch_size=8, num_workers=8, pin_memory=True, shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2021-11-04T18:47:04.694245Z","iopub.execute_input":"2021-11-04T18:47:04.695020Z","iopub.status.idle":"2021-11-04T18:47:04.781127Z","shell.execute_reply.started":"2021-11-04T18:47:04.694986Z","shell.execute_reply":"2021-11-04T18:47:04.780351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Training**","metadata":{"papermill":{"duration":0.014232,"end_time":"2021-10-26T02:08:43.77145","exception":false,"start_time":"2021-10-26T02:08:43.757218","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# !!! Run more epochs :)\nLEARNING_RATE = 5e-4\nEPOCHS = 3","metadata":{"execution":{"iopub.status.busy":"2021-11-04T18:47:06.545228Z","iopub.execute_input":"2021-11-04T18:47:06.545484Z","iopub.status.idle":"2021-11-04T18:47:06.550120Z","shell.execute_reply.started":"2021-11-04T18:47:06.545456Z","shell.execute_reply":"2021-11-04T18:47:06.549426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.cuda()","metadata":{"execution":{"iopub.status.busy":"2021-11-04T18:47:07.316657Z","iopub.execute_input":"2021-11-04T18:47:07.317301Z","iopub.status.idle":"2021-11-04T18:47:10.158091Z","shell.execute_reply.started":"2021-11-04T18:47:07.317268Z","shell.execute_reply":"2021-11-04T18:47:10.157396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_samples_val = len(ds_val)\nn_batches_val = len(ds_val)\nn_batches_train = len(dl_train)\nn_batches_train_extra = len(dl_train_extra)\ncriterion = CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)","metadata":{"execution":{"iopub.status.busy":"2021-11-04T18:47:10.159681Z","iopub.execute_input":"2021-11-04T18:47:10.160083Z","iopub.status.idle":"2021-11-04T18:47:10.166914Z","shell.execute_reply.started":"2021-11-04T18:47:10.160044Z","shell.execute_reply":"2021-11-04T18:47:10.166152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(1, EPOCHS + 1):\n    print(f\"Starting epoch: {epoch} / {EPOCHS}\")\n    \n    train_loss = 0.0\n    train_extra_loss = 0.0\n    optimizer.zero_grad()\n    model.train()\n    \n    # Train on extra data\n    for batch_idx, batch in enumerate(dl_train_extra):\n        \n        # Predict\n        images, labels = batch['image'], batch['label']\n        images, labels = images.cuda(),  labels.cuda()\n        \n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        \n        # Back prop\n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n        train_extra_loss += loss.item()\n    \n    # Train on train data\n    for batch_idx, batch in enumerate(dl_train):\n        \n        # Predict\n        images, labels = batch['image'], batch['label']\n        images, labels = images.cuda(),  labels.cuda()\n        \n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        \n        # Back prop\n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n        train_loss += loss.item()\n    \n    # Validate\n    model.eval()\n    loss = 0\n    correct = 0\n    \n    with torch.no_grad():\n        for batch_idx, batch in enumerate(dl_val, 1):\n            images, labels = batch['image'], batch['label']\n            images, labels = images.cuda(),  labels.cuda()\n            preds = model(images)\n            final_pred = preds.argmax(dim=1)\n            correct += (final_pred == labels).sum().item()\n            loss += criterion(preds, labels)\n\n    train_loss = train_loss / n_batches_train\n    train_extra_loss = train_extra_loss / n_batches_train_extra\n    loss = loss / n_batches_val\n    acc = correct / n_samples_val\n    \n    print(f\"Epoch: {epoch} - Train Extra Loss {train_extra_loss:.5f}. Train Loss {train_loss:.5f}. Val. Loss: {loss:.5f} Accuracy: {acc*100:.4f}%\")","metadata":{"papermill":{"duration":198.280697,"end_time":"2021-10-26T02:12:02.066195","exception":false,"start_time":"2021-10-26T02:08:43.785498","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-11-04T18:47:12.424547Z","iopub.execute_input":"2021-11-04T18:47:12.425073Z","iopub.status.idle":"2021-11-04T18:49:00.299867Z","shell.execute_reply.started":"2021-11-04T18:47:12.425033Z","shell.execute_reply":"2021-11-04T18:49:00.298169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model, 'cell-classification-model.pth')","metadata":{"papermill":{"duration":0.171381,"end_time":"2021-10-26T02:12:02.254084","exception":false,"start_time":"2021-10-26T02:12:02.082703","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2021-11-04T18:46:19.176201Z","iopub.status.idle":"2021-11-04T18:46:19.176972Z","shell.execute_reply.started":"2021-11-04T18:46:19.176713Z","shell.execute_reply":"2021-11-04T18:46:19.176754Z"},"trusted":true},"execution_count":null,"outputs":[]}]}