{"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-30T10:33:05.978261Z","iopub.execute_input":"2022-10-30T10:33:05.978943Z","iopub.status.idle":"2022-10-30T10:33:33.049323Z","shell.execute_reply.started":"2022-10-30T10:33:05.978908Z","shell.execute_reply":"2022-10-30T10:33:33.048104Z"},"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-30T10:33:33.054886Z","iopub.execute_input":"2022-10-30T10:33:33.057095Z","iopub.status.idle":"2022-10-30T10:33:40.72462Z","shell.execute_reply.started":"2022-10-30T10:33:33.057054Z","shell.execute_reply":"2022-10-30T10:33:40.723409Z"},"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/segmentations_npy/*\")))\n\npaths[:5]","metadata":{"execution":{"iopub.status.busy":"2022-10-30T10:44:31.354725Z","iopub.execute_input":"2022-10-30T10:44:31.355095Z","iopub.status.idle":"2022-10-30T10:44:31.364635Z","shell.execute_reply.started":"2022-10-30T10:44:31.355064Z","shell.execute_reply":"2022-10-30T10:44:31.363531Z"},"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-30T10:42:51.941065Z","iopub.execute_input":"2022-10-30T10:42:51.941847Z","iopub.status.idle":"2022-10-30T10:42:51.946745Z","shell.execute_reply.started":"2022-10-30T10:42:51.941808Z","shell.execute_reply":"2022-10-30T10:42:51.945626Z"},"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-30T10:42:52.2725Z","iopub.execute_input":"2022-10-30T10:42:52.273232Z","iopub.status.idle":"2022-10-30T10:42:52.28105Z","shell.execute_reply.started":"2022-10-30T10:42:52.273186Z","shell.execute_reply":"2022-10-30T10:42:52.280098Z"},"trusted":true},"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            try:\n                mask = np.load(path)\n            except:\n                mask = np.load(path.replace(\"segmentations_npy\", \"prediction_sagittal_b1v1\").replace('segmentations-npy', 'prediction-sagittal-b1v1'))\n                \n            image = np.load(path.replace(\"segmentations_npy\", \"train_sagittal\").replace('segmentations-npy', 'train-sagittal'))\n            \n            #image = image[:, :, image.shape[2]//2]\n            try:\n                mask = mask[:, :, mask.shape[2]//2]\n            except:\n                pass\n            \n            image = np.stack([image]*3, -1)\n            \n            #mask_ = np.zeros((mask.shape[0], mask.shape[1], 1), dtype=np.float32)\n            #for u in np.unique(mask):\n            #    if u: mask_[:, :, np.clip(u-1, 0, 0)][mask==u] = 1.\n        \n            #mask = mask_\n            \n            mask = np.expand_dims(np.clip(mask, 0, 1), -1)\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-30T10:57:20.814079Z","iopub.execute_input":"2022-10-30T10:57:20.815364Z","iopub.status.idle":"2022-10-30T10:57:20.827914Z","shell.execute_reply.started":"2022-10-30T10:57:20.8153Z","shell.execute_reply":"2022-10-30T10:57:20.826543Z"},"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 = np.array(paths[folds[CFG.FOLD][0]].tolist() * 10)\n    valid_paths = paths[folds[CFG.FOLD][1]]\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=0, 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-30T10:52:56.280097Z","iopub.execute_input":"2022-10-30T10:52:56.281047Z","iopub.status.idle":"2022-10-30T10:52:56.762964Z","shell.execute_reply.started":"2022-10-30T10:52:56.280997Z","shell.execute_reply":"2022-10-30T10:52:56.761969Z"},"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=1,)\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 = 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-30T10:47:01.338139Z","iopub.execute_input":"2022-10-30T10:47:01.339095Z","iopub.status.idle":"2022-10-30T10:47:01.352554Z","shell.execute_reply.started":"2022-10-30T10:47:01.339051Z","shell.execute_reply":"2022-10-30T10:47:01.351666Z"},"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-30T10:48:15.715359Z","iopub.execute_input":"2022-10-30T10:48:15.716317Z","iopub.status.idle":"2022-10-30T10:48:58.586857Z","shell.execute_reply.started":"2022-10-30T10:48:15.716266Z","shell.execute_reply":"2022-10-30T10:48:58.584838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DIR = \"/kaggle/input/\"\n\npaths = np.array(sorted(glob(f\"{DIR}/rsna-2022-segmentations-npy/segmentations_npy/*\")))\n\ntrain = pd.read_csv('/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv')\n\npseudo = np.array([f'{DIR}/rsna-2022-segmentations-npy/segmentations_npy//{p}.npy' for p in train.StudyInstanceUID.values if f\"{p}.npy\" not in os.listdir(f\"{DIR}/rsna-2022-segmentations-npy/segmentations_npy/\")])\n\npaths[:5], pseudo.shape","metadata":{"execution":{"iopub.status.busy":"2022-10-30T10:56:57.830013Z","iopub.execute_input":"2022-10-30T10:56:57.831276Z","iopub.status.idle":"2022-10-30T10:56:58.664263Z","shell.execute_reply.started":"2022-10-30T10:56:57.83121Z","shell.execute_reply":"2022-10-30T10:56:58.663337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CFG.EPOCHS = 1 if DEBUG else 24\nCFG.LR = 1e-3\nCFG.WARMUP_EPOCHS = 4\nCFG.WARMUP_LR = 1e-6\nCFG.V = \"3\"","metadata":{"execution":{"iopub.status.busy":"2022-10-30T10:56:59.312643Z","iopub.execute_input":"2022-10-30T10:56:59.313005Z","iopub.status.idle":"2022-10-30T10:56:59.318232Z","shell.execute_reply.started":"2022-10-30T10:56:59.312974Z","shell.execute_reply":"2022-10-30T10:56:59.317166Z"},"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    train_paths = np.array(train_paths.tolist() + pseudo.tolist())\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=0, 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-30T10:57:32.720748Z","iopub.execute_input":"2022-10-30T10:57:32.721145Z","iopub.status.idle":"2022-10-30T10:57:33.230096Z","shell.execute_reply.started":"2022-10-30T10:57:32.721089Z","shell.execute_reply":"2022-10-30T10:57:33.229138Z"},"trusted":true},"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-30T10:58:08.069401Z","iopub.execute_input":"2022-10-30T10:58:08.069783Z","iopub.status.idle":"2022-10-30T10:59:46.145527Z","shell.execute_reply.started":"2022-10-30T10:58:08.06975Z","shell.execute_reply":"2022-10-30T10:59:46.143414Z"},"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            try:\n                mask = np.load(path)\n            except:\n                mask = np.load(path.replace(\"segmentations_npy\", \"prediction_sagittal_b1v1\").replace('segmentations-npy', 'prediction-sagittal-b1v1'))\n                \n            image = np.load(path.replace(\"segmentations_npy\", \"train_sagittal\").replace('segmentations-npy', 'train-sagittal'))\n            \n            #image = image[:, :, image.shape[2]//2]\n            try:\n                mask = mask[:, :, mask.shape[2]//2]\n            except:\n                pass\n            \n            image = np.stack([image]*1, -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: mask_[:, :, np.clip(u-1, 0, 7)][mask==u] = 1.\n                    \n            mask = mask_\n            \n            #mask = np.expand_dims(np.clip(mask, 0, 1), -1)\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-30T11:14:49.702728Z","iopub.execute_input":"2022-10-30T11:14:49.703208Z","iopub.status.idle":"2022-10-30T11:14:49.715827Z","shell.execute_reply.started":"2022-10-30T11:14:49.703169Z","shell.execute_reply":"2022-10-30T11:14:49.714848Z"},"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    train_paths = np.array(train_paths.tolist() + pseudo.tolist())\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.25, limit=(25, -25)),\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-30T11:14:49.921261Z","iopub.execute_input":"2022-10-30T11:14:49.922234Z","iopub.status.idle":"2022-10-30T11:14:51.794202Z","shell.execute_reply.started":"2022-10-30T11:14:49.922197Z","shell.execute_reply":"2022-10-30T11:14:51.793095Z"},"trusted":true},"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=1, 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 = 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-30T11:14:51.796893Z","iopub.execute_input":"2022-10-30T11:14:51.797746Z","iopub.status.idle":"2022-10-30T11:14:51.811105Z","shell.execute_reply.started":"2022-10-30T11:14:51.7977Z","shell.execute_reply":"2022-10-30T11:14:51.809893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CFG.EPOCHS = 1 if DEBUG else 24\nCFG.LR = 1e-3\nCFG.WARMUP_EPOCHS = 4\nCFG.WARMUP_LR = 1e-6\nCFG.V = \"10\"","metadata":{"execution":{"iopub.status.busy":"2022-10-30T11:14:53.8604Z","iopub.execute_input":"2022-10-30T11:14:53.860766Z","iopub.status.idle":"2022-10-30T11:14:53.866064Z","shell.execute_reply.started":"2022-10-30T11:14:53.860733Z","shell.execute_reply":"2022-10-30T11:14:53.864574Z"},"trusted":true},"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\"/kaggle/input/rsna-2022-seg-b1-v3/f0.ckpt\")['state_dict']\n        LEFTOUTS = ['feature_extractor.segmentation_head.0.weight', 'feature_extractor.segmentation_head.0.bias']\n        st = {k:st[k] for k in st if k not in LEFTOUTS}\n        model.load_state_dict(st, strict=False)\n        \n        #for p in model.feature_extractor.encoder.parameters(): p.requires_grad = False\n        #for p in model.feature_extractor.decoder.parameters(): p.requires_grad = False\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-30T11:18:45.048918Z","iopub.execute_input":"2022-10-30T11:18:45.049347Z","iopub.status.idle":"2022-10-30T11:19:56.987128Z","shell.execute_reply.started":"2022-10-30T11:18:45.049308Z","shell.execute_reply":"2022-10-30T11:19:56.97963Z"},"trusted":true},"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=1, classes=8,)\n        \n        self.sigmoid = nn.Sigmoid()\n        \n    def forward(self, inp):\n        masks = self.feature_extractor(inp)\n        return masks","metadata":{"execution":{"iopub.status.busy":"2022-10-30T11:22:00.616055Z","iopub.execute_input":"2022-10-30T11:22:00.617155Z","iopub.status.idle":"2022-10-30T11:22:00.623437Z","shell.execute_reply.started":"2022-10-30T11:22:00.617082Z","shell.execute_reply":"2022-10-30T11:22:00.622367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nfor _ in range(5):\n    model = Model()\n    model.eval()\n    model.cuda()\n    st = torch.load(f'/kaggle/input/try2-seg-b1v10-sagview-full/f{_}.ckpt')['state_dict']\n    model.load_state_dict(st)\n    models.append(copy.deepcopy(model))","metadata":{"execution":{"iopub.status.busy":"2022-10-30T11:22:25.634307Z","iopub.execute_input":"2022-10-30T11:22:25.634674Z","iopub.status.idle":"2022-10-30T11:22:34.156323Z","shell.execute_reply.started":"2022-10-30T11:22:25.634642Z","shell.execute_reply":"2022-10-30T11:22:34.155208Z"},"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":"paths = np.array(glob(f\"/kaggle/input/rsna-2022-train-sagittal/train_sagittal/*\"))\npaths[:5], paths.shape","metadata":{"execution":{"iopub.status.busy":"2022-10-30T11:24:57.769062Z","iopub.execute_input":"2022-10-30T11:24:57.770148Z","iopub.status.idle":"2022-10-30T11:24:57.785612Z","shell.execute_reply.started":"2022-10-30T11:24:57.770069Z","shell.execute_reply":"2022-10-30T11:24:57.784513Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir /kaggle/working/classes_volumes_b1v10/\n\nfor path in tqdm(paths):\n    sag = np.load(path)\n    \n    sh = list(sag.shape[:2])\n    sh[0], sh[1] = sh[1], sh[0]\n    \n    image = torch.as_tensor(cv2.resize(sag, (256, 256))).unsqueeze(0).unsqueeze(0).float()/255\n    \n    with torch.no_grad():\n        outputs = []\n        for model in models:\n            output = model.sigmoid(model(image.cuda())).detach().cpu().numpy()[0].transpose(1, 2, 0)\n            outputs.append(output)\n        output = np.mean(outputs, 0)\n        output = cv2.resize(output, sh)\n        output[output>0.3] = 1\n        output[output<0.3] = 0\n    \n    preds = []\n    for _ in output:\n        classes = np.sum(_, 0)\n        if np.any(classes):\n            preds.append(np.argmax(classes)+1)\n        else:\n            preds.append(100)\n    \n    np.save(f\"/kaggle/working/classes_volumes_b1v10/{path.split('/')[-1]}\", preds)\n    \n    #break","metadata":{"execution":{"iopub.status.busy":"2022-10-30T11:27:04.720589Z","iopub.execute_input":"2022-10-30T11:27:04.720982Z","iopub.status.idle":"2022-10-30T11:27:05.915338Z","shell.execute_reply.started":"2022-10-30T11:27:04.720943Z","shell.execute_reply":"2022-10-30T11:27:05.91392Z"},"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":[]}]}