{"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 -q segmentation_models_pytorch\n!pip install -q pytorch-lightning-bolts","metadata":{"execution":{"iopub.status.busy":"2022-10-31T08:15:52.808017Z","iopub.execute_input":"2022-10-31T08:15:52.808462Z","iopub.status.idle":"2022-10-31T08:16:18.607266Z","shell.execute_reply.started":"2022-10-31T08:15:52.808381Z","shell.execute_reply":"2022-10-31T08:16:18.606051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport os\nfrom glob import glob\nimport copy\nimport time\nimport math\n\nimport cv2\nimport matplotlib.pyplot as plt\nfrom skimage import img_as_ubyte\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom sklearn.model_selection import *\nfrom sklearn.metrics import *\n\nimport torch\nfrom torch import nn, optim\nfrom torch.utils.data import Dataset, DataLoader\nimport pytorch_lightning as pl\nimport pl_bolts as pb\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2022-10-31T08:16:18.610214Z","iopub.execute_input":"2022-10-31T08:16:18.610861Z","iopub.status.idle":"2022-10-31T08:16:25.575643Z","shell.execute_reply.started":"2022-10-31T08:16:18.610813Z","shell.execute_reply":"2022-10-31T08:16:25.574517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DIR = \"/kaggle/input/\"\n\npaths = np.array(sorted(glob(f\"{DIR}/rsna-2022-segmentations-npy-slices/segmentations_npy_slices/*\")))\n\npaths[:5], paths.shape","metadata":{"execution":{"iopub.status.busy":"2022-10-31T08:16:25.577341Z","iopub.execute_input":"2022-10-31T08:16:25.57799Z","iopub.status.idle":"2022-10-31T08:16:26.219421Z","shell.execute_reply.started":"2022-10-31T08:16:25.577949Z","shell.execute_reply":"2022-10-31T08:16:26.218546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = 1 # CHANGE THIS TO 0 IF WANT TO ACTUALLY TRAIN","metadata":{"execution":{"iopub.status.busy":"2022-10-31T08:16:28.246685Z","iopub.execute_input":"2022-10-31T08:16:28.247043Z","iopub.status.idle":"2022-10-31T08:16:28.252065Z","shell.execute_reply.started":"2022-10-31T08:16:28.247012Z","shell.execute_reply":"2022-10-31T08:16:28.250791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    SEED = 42\n    SPLITS = 5\n    FOLD = 0\n    \n    SZ_H = 256\n    SZ_W = 256\n    \n    TRN_BS = 8\n    VAL_BS = 8\n    ACCUMS = 4\n    \n    ACCL = None#\"dp\"\n    \n    EPOCHS = 1 if DEBUG else 96\n    LR = 1e-3\n    WARMUP_EPOCHS = 24\n    WARMUP_LR = 1e-6\n    \n    NAME = \"b1\"\n    V = \"1\"\n    \npl.seed_everything(CFG.SEED)\nOUTPUT_FOLDER = f\"/kaggle/working/{CFG.NAME}_v{CFG.V}/\"","metadata":{"execution":{"iopub.status.busy":"2022-10-31T08:16:30.941488Z","iopub.execute_input":"2022-10-31T08:16:30.941981Z","iopub.status.idle":"2022-10-31T08:16:30.957078Z","shell.execute_reply.started":"2022-10-31T08:16:30.941931Z","shell.execute_reply":"2022-10-31T08:16:30.95598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GiTractDataset(Dataset):\n    def __init__(self, paths, transforms=None):\n        self.paths = paths\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.paths)\n    \n    def __getitem__(self, i):\n        try:\n            path = self.paths[i]\n        \n            mask = np.load(path)\n            image = np.load(path.replace('segmentations_npy_slices', 'segmentation_images_kaggle').replace('rsna-2022-segmentations-npy-slices', 'segmentation-images-kaggle'))\n        \n            image = np.stack([image]*3, -1)\n        \n            mask_ = np.zeros((mask.shape[0], mask.shape[1], 8), dtype=np.float32)\n            for u in np.unique(mask):\n                if u:\n                    mask_[:, :, np.clip(u-1, 0, 7)][mask==u] = 1.\n        \n            mask = mask_\n        \n            if self.transforms:\n                transformed = self.transforms(image=image, mask=mask)\n                image = transformed['image']\n                mask = transformed['mask']\n                mask = torch.as_tensor(mask.numpy().transpose(2, 0, 1))\n                \n                if image.dtype==torch.uint8: image = image.float()/255\n        \n            return image, mask\n        \n        except:\n            return torch.zeros((3, CFG.SZ_H, CFG.SZ_W)).float(), torch.zeros((8, CFG.SZ_H, CFG.SZ_W)).float()","metadata":{"execution":{"iopub.status.busy":"2022-10-31T08:16:31.631591Z","iopub.execute_input":"2022-10-31T08:16:31.631959Z","iopub.status.idle":"2022-10-31T08:16:31.643793Z","shell.execute_reply.started":"2022-10-31T08:16:31.631927Z","shell.execute_reply":"2022-10-31T08:16:31.642816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"folds = [*GroupKFold(n_splits=CFG.SPLITS).split(paths, groups=[x.split('.')[-2].split('_')[0] for x in paths])]\n\ndef get_loaders():\n    \n    train_paths = paths[folds[CFG.FOLD][0]]\n    valid_paths = paths[folds[CFG.FOLD][1]]\n    \n    if DEBUG:\n        train_paths = train_paths[:1000]\n        valid_paths = train_paths[:1000]\n    \n    train_augs = A.Compose([\n        A.Resize(CFG.SZ_H, CFG.SZ_W),\n        #A.RandomResizedCrop(CFG.SZ_H, CFG.SZ_W, ratio=[0.9, 1.1], scale=[0.9, 1.1]),\n        #A.OneOf([\n        #    A.Resize(CFG.SZ_H, CFG.SZ_W),\n        #    A.RandomResizedCrop(CFG.SZ_H, CFG.SZ_W, ratio=[0.6, 1.4], scale=[0.5, 1.5]),\n        #], p=1.),\n        A.Perspective(p=0.5),\n        #A.HorizontalFlip(p=0.25),\n        #A.VerticalFlip(p=0.25),\n        #A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, p=.5),\n        A.Rotate(p=0.5, limit=(45, -45)),\n        #A.RandomContrast(limit=(0.5, 0.5), p=1.),\n        #A.RandomBrightnessContrast(p=0.5),\n        #A.Cutout(p=0.25, max_h_size=CFG.SZ//4, max_w_size=CFG.SZ//4, num_holes=4),\n        #A.Normalize(),\n        ToTensorV2(),\n    ])\n    \n    valid_augs = A.Compose([\n        A.Resize(CFG.SZ_H, CFG.SZ_W),\n        #A.RandomContrast(limit=(0.2, 0.2), p=1.),\n        #A.Normalize(),\n        ToTensorV2()\n    ])\n    \n    train_dataset = GiTractDataset(train_paths, train_augs)\n    valid_dataset = GiTractDataset(valid_paths, valid_augs)\n    \n    train_loader = DataLoader(train_dataset, batch_size=CFG.TRN_BS, shuffle=True, num_workers=8, pin_memory=False)\n    valid_loader = DataLoader(valid_dataset, batch_size=CFG.VAL_BS, shuffle=False, num_workers=8, pin_memory=False)\n    \n    return train_loader, valid_loader#, train_data, valid_data\n\ntrain_loader, valid_loader = get_loaders()\nfor d in valid_loader: break\nplt.imshow(d[0][0].numpy().transpose(1, 2, 0))","metadata":{"execution":{"iopub.status.busy":"2022-10-31T08:18:47.573541Z","iopub.execute_input":"2022-10-31T08:18:47.573952Z","iopub.status.idle":"2022-10-31T08:18:50.633372Z","shell.execute_reply.started":"2022-10-31T08:18:47.573912Z","shell.execute_reply":"2022-10-31T08:18:50.630644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Model(pl.LightningModule):\n    def __init__(self):\n        super(Model, self).__init__()\n        #tf_efficientnet_b0_ns resnest50d_4s2x40d seresnext50_32x4d tf_efficientnetv2_m_in21ft1k\n        self.feature_extractor = smp.Unet('tu-tf_efficientnet_b1_ns', in_channels=3, classes=8, aux_params={'classes': 8},)\n        \n        self.flatten = nn.Flatten()\n        self.sigmoid = nn.Sigmoid()\n        self.softmax = nn.Softmax(-1)\n        self.bce = nn.BCEWithLogitsLoss()\n        self.dice = smp.losses.DiceLoss(mode=smp.losses.MULTILABEL_MODE)\n        \n    def forward(self, inp):\n        masks, logits = self.feature_extractor(inp)\n        return masks\n    \n    def _criterion(self, outputs, targets):\n        #bce = self.bce(outputs, targets)\n        dice = self.dice(outputs, targets)\n        loss = dice# + bce\n        return loss\n    \n    def _validation_score(self, outputs, targets):\n        dice = self.dice(outputs, targets)\n        return 1-dice\n    \n    def training_step(self, batch, idx):\n        inputs, masks = batch\n        \n        if len(inputs)>1:\n            outputs = self(inputs)\n        else:\n            outputs = self(torch.cat([inputs, inputs]))[:1]\n            \n        loss = self._criterion(outputs, masks)\n        \n        return loss\n    \n    def validation_step(self, batch, idx):\n        inputs, targets = batch\n        \n        outputs = self(inputs)\n        \n        score = self._validation_score(outputs, targets)\n        \n        self.log(\"m\", score, on_epoch=True, prog_bar=True, sync_dist=True)\n        \n    def configure_optimizers(self):\n        optimizer = optim.AdamW(self.parameters(), lr=CFG.LR, weight_decay=1e-5)\n        #optimizer = optim.SGD(self.parameters(), lr=CFG.LR, weight_decay=1e-5)\n        #optimizer = AdamSGDWeighted(self.parameters(), lr=CFG.LR, weight_decay=1e-5, adam_w=0.4, sgd_w=0.6)\n        \n        \n        scheduler = pb.optimizers.lr_scheduler.LinearWarmupCosineAnnealingLR(optimizer, \n                                                                             warmup_epochs=CFG.WARMUP_EPOCHS, \n                                                                             max_epochs=CFG.EPOCHS,\n                                                                             warmup_start_lr=CFG.WARMUP_LR)\n        \n        return [optimizer], [scheduler]\n    \n    def get_progress_bar_dict(self):\n        tqdm_dict = super().get_progress_bar_dict()\n        if 'v_num' in tqdm_dict: del tqdm_dict['v_num']\n        return tqdm_dict","metadata":{"execution":{"iopub.status.busy":"2022-10-31T08:18:59.737122Z","iopub.execute_input":"2022-10-31T08:18:59.738249Z","iopub.status.idle":"2022-10-31T08:18:59.752292Z","shell.execute_reply.started":"2022-10-31T08:18:59.738189Z","shell.execute_reply":"2022-10-31T08:18:59.750998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs(OUTPUT_FOLDER, exist_ok=1)\n\nwith open(f\"{OUTPUT_FOLDER}/scores_f{CFG.FOLD}.txt\", 'w+') as f:\n\n    for F in range(CFG.FOLD, CFG.SPLITS):\n    #for F in range(CFG.FOLD, CFG.FOLD+1):\n        print(f\"FOLD {F}\")\n\n        CFG.FOLD = F\n        train_loader, valid_loader = get_loaders()\n        \n        checkpoint_callback = pl.callbacks.ModelCheckpoint(\n        monitor=\"m\",\n        dirpath=f\"{OUTPUT_FOLDER}\",\n        filename=f\"f{F}\",\n        save_top_k=1,\n        mode=\"max\",\n        )\n        \n        logger = pl.loggers.CSVLogger(f\"{OUTPUT_FOLDER}\", name=\"log\")\n        \n        warnings.filterwarnings(\"ignore\")\n        \n        trainer = pl.Trainer(gpus=-1, accelerator=CFG.ACCL,\n                             accumulate_grad_batches=CFG.ACCUMS, \n                             deterministic=True, precision=16,\n                             max_epochs=CFG.EPOCHS,\n                             callbacks=[checkpoint_callback],\n                             logger=logger)\n        \n        model = Model()\n        #st = torch.load(f\"/mnt/md0/gi_tract_seg/AAA/TRY3_GOOD_SEGMENTATION/b4_v4/f0.pt\")\n        #LEFTOUTS = ['feature_extractor.segmentation_head.0.weight', 'feature_extractor.segmentation_head.0.bias',\n        #            'feature_extractor.classification_head.3.weight', 'feature_extractor.classification_head.3.bias']\n        #st = {k:st[k] for k in st if k not in LEFTOUTS}\n        #model.load_state_dict(st, strict=False)\n        #for param in model.feature_extractor.parameters(): param.requires_grad = False\n        #model.feature_extractor2.load_state_dict(model.feature_extractor.state_dict())\n        \n        trainer.fit(model, train_loader, valid_loader)\n        \n        time.sleep(10)\n        \n        try:\n            st = torch.load(f\"{OUTPUT_FOLDER}/f{F}.ckpt\")['state_dict']\n            model.load_state_dict(st)\n            torch.save(st, f\"{OUTPUT_FOLDER}/f{F}.pt\")\n        except: print(\"ERROR LOADING\")\n        #torch.jit.save(torch.jit.script(model), f\"{DIR}/{NAME}_v{V}/f{F}.pt\")\n        \n        '''\n        m = torch.jit.load(f\"{DIR}/{NAME}_v{V}/f{F}.pt\")\n        m.eval()\n        m(torch.zeros((2, 3, 512, 512)).unsqueeze(0))\n        '''\n        \n        print(checkpoint_callback.best_model_score)\n        \n        f.write(f\"{checkpoint_callback.best_model_score} \\n\")","metadata":{"execution":{"iopub.status.busy":"2022-10-31T08:19:00.707467Z","iopub.execute_input":"2022-10-31T08:19:00.707839Z","iopub.status.idle":"2022-10-31T08:20:06.963814Z","shell.execute_reply.started":"2022-10-31T08:19:00.707807Z","shell.execute_reply":"2022-10-31T08:20:06.961767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}