{"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\n!pip install pytorch-lightning-bolts","metadata":{"execution":{"iopub.status.busy":"2023-01-14T05:10:59.443611Z","iopub.execute_input":"2023-01-14T05:10:59.444506Z","iopub.status.idle":"2023-01-14T05:11:26.489597Z","shell.execute_reply.started":"2023-01-14T05:10:59.444391Z","shell.execute_reply":"2023-01-14T05:11:26.488399Z"},"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":"2023-01-14T05:11:26.493332Z","iopub.execute_input":"2023-01-14T05:11:26.493663Z","iopub.status.idle":"2023-01-14T05:11:33.603238Z","shell.execute_reply.started":"2023-01-14T05:11:26.493629Z","shell.execute_reply":"2023-01-14T05:11:33.602091Z"},"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":"2023-01-14T05:11:33.605267Z","iopub.execute_input":"2023-01-14T05:11:33.605926Z","iopub.status.idle":"2023-01-14T05:11:33.661689Z","shell.execute_reply.started":"2023-01-14T05:11:33.605884Z","shell.execute_reply":"2023-01-14T05:11:33.660648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEBUG = 0 # CHANGE THIS TO 0 IF WANT TO ACTUALLY TRAIN","metadata":{"execution":{"iopub.status.busy":"2023-01-14T05:11:33.664735Z","iopub.execute_input":"2023-01-14T05:11:33.665114Z","iopub.status.idle":"2023-01-14T05:11:33.669566Z","shell.execute_reply.started":"2023-01-14T05:11:33.665077Z","shell.execute_reply":"2023-01-14T05:11:33.66843Z"},"trusted":true},"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":"2023-01-14T05:11:33.671135Z","iopub.execute_input":"2023-01-14T05:11:33.671825Z","iopub.status.idle":"2023-01-14T05:11:33.683388Z","shell.execute_reply.started":"2023-01-14T05:11:33.67179Z","shell.execute_reply":"2023-01-14T05:11:33.682476Z"},"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":"2023-01-14T05:11:33.684983Z","iopub.execute_input":"2023-01-14T05:11:33.685391Z","iopub.status.idle":"2023-01-14T05:11:33.69644Z","shell.execute_reply.started":"2023-01-14T05:11:33.685357Z","shell.execute_reply":"2023-01-14T05:11:33.69521Z"},"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":"2023-01-14T05:11:33.697996Z","iopub.execute_input":"2023-01-14T05:11:33.698371Z","iopub.status.idle":"2023-01-14T05:11:39.308181Z","shell.execute_reply.started":"2023-01-14T05:11:33.698337Z","shell.execute_reply":"2023-01-14T05:11:39.307274Z"},"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=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":"2023-01-14T05:11:39.309298Z","iopub.execute_input":"2023-01-14T05:11:39.310342Z","iopub.status.idle":"2023-01-14T05:11:39.32805Z","shell.execute_reply.started":"2023-01-14T05:11:39.310302Z","shell.execute_reply":"2023-01-14T05:11:39.327117Z"},"trusted":true},"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        \n        model = torch.quantization.convert(model.eval(), inplace=False)\n    \n        optimizer = torch.optim.SGD(model.parameters(), lr = 1e-3)\n        model.qconfig = torch.ao.quantization.get_default_qat_qconfig('fbgemm')\n    \n    \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":"2023-01-14T05:11:39.329509Z","iopub.execute_input":"2023-01-14T05:11:39.330074Z","iopub.status.idle":"2023-01-14T05:59:28.376511Z","shell.execute_reply.started":"2023-01-14T05:11:39.330038Z","shell.execute_reply":"2023-01-14T05:59:28.375359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"folds","metadata":{"execution":{"iopub.status.busy":"2023-01-14T05:59:28.381303Z","iopub.execute_input":"2023-01-14T05:59:28.381619Z","iopub.status.idle":"2023-01-14T05:59:28.391134Z","shell.execute_reply.started":"2023-01-14T05:59:28.381587Z","shell.execute_reply":"2023-01-14T05:59:28.390104Z"},"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\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":"2023-01-14T05:59:28.39245Z","iopub.execute_input":"2023-01-14T05:59:28.393149Z","iopub.status.idle":"2023-01-14T05:59:29.203521Z","shell.execute_reply.started":"2023-01-14T05:59:28.393113Z","shell.execute_reply":"2023-01-14T05:59:29.20255Z"},"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":"2023-01-14T05:59:29.204803Z","iopub.execute_input":"2023-01-14T05:59:29.205181Z","iopub.status.idle":"2023-01-14T05:59:29.212903Z","shell.execute_reply.started":"2023-01-14T05:59:29.205144Z","shell.execute_reply":"2023-01-14T05:59:29.211956Z"},"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":"2023-01-14T05:59:29.214446Z","iopub.execute_input":"2023-01-14T05:59:29.21512Z","iopub.status.idle":"2023-01-14T05:59:29.681408Z","shell.execute_reply.started":"2023-01-14T05:59:29.215083Z","shell.execute_reply":"2023-01-14T05:59:29.680314Z"},"trusted":true},"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        \n        model = torch.quantization.convert(model.eval(), inplace=False)\n    \n        optimizer = torch.optim.SGD(model.parameters(), lr = 1e-3)\n        model.qconfig = torch.ao.quantization.get_default_qat_qconfig('fbgemm')\n        \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":"2023-01-14T05:59:29.683103Z","iopub.execute_input":"2023-01-14T05:59:29.683802Z","iopub.status.idle":"2023-01-14T06:15:59.959501Z","shell.execute_reply.started":"2023-01-14T05:59:29.683761Z","shell.execute_reply":"2023-01-14T06:15:59.958057Z"},"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]*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":"2023-01-14T06:15:59.964867Z","iopub.execute_input":"2023-01-14T06:15:59.965638Z","iopub.status.idle":"2023-01-14T06:15:59.98835Z","shell.execute_reply.started":"2023-01-14T06:15:59.96556Z","shell.execute_reply":"2023-01-14T06:15:59.987362Z"},"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":"2023-01-14T06:15:59.989914Z","iopub.execute_input":"2023-01-14T06:15:59.990544Z","iopub.status.idle":"2023-01-14T06:16:01.754439Z","shell.execute_reply.started":"2023-01-14T06:15:59.990505Z","shell.execute_reply":"2023-01-14T06:16:01.75334Z"},"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":"2023-01-14T06:16:01.756479Z","iopub.execute_input":"2023-01-14T06:16:01.757176Z","iopub.status.idle":"2023-01-14T06:16:01.770535Z","shell.execute_reply.started":"2023-01-14T06:16:01.757127Z","shell.execute_reply":"2023-01-14T06:16:01.769467Z"},"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":"2023-01-14T06:16:01.771924Z","iopub.execute_input":"2023-01-14T06:16:01.772356Z","iopub.status.idle":"2023-01-14T06:16:01.786537Z","shell.execute_reply.started":"2023-01-14T06:16:01.77232Z","shell.execute_reply":"2023-01-14T06:16:01.785604Z"},"trusted":true},"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        \n        model = torch.quantization.convert(model.eval(), inplace=False)\n    \n        optimizer = torch.optim.SGD(model.parameters(), lr = 1e-3)\n        model.qconfig = torch.ao.quantization.get_default_qat_qconfig('fbgemm')\n        \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":"2023-01-14T06:16:01.788221Z","iopub.execute_input":"2023-01-14T06:16:01.788796Z","iopub.status.idle":"2023-01-14T08:00:38.353612Z","shell.execute_reply.started":"2023-01-14T06:16:01.788757Z","shell.execute_reply":"2023-01-14T08:00:38.349002Z"},"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.sigmoid = nn.Sigmoid()\n        \n    def forward(self, inp):\n        masks = self.feature_extractor(inp)\n        return masks","metadata":{"execution":{"iopub.status.busy":"2023-01-14T08:00:38.359518Z","iopub.execute_input":"2023-01-14T08:00:38.36147Z","iopub.status.idle":"2023-01-14T08:00:38.382207Z","shell.execute_reply.started":"2023-01-14T08:00:38.36139Z","shell.execute_reply":"2023-01-14T08:00:38.380713Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nfor _ in range(5):\n    model = Model()\n    model = torch.quantization.convert(model.eval(), inplace=False)\n    \n    optimizer = torch.optim.SGD(model.parameters(), lr = 1e-3)\n    model.qconfig = torch.ao.quantization.get_default_qat_qconfig('fbgemm')\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":"2023-01-14T08:00:38.384342Z","iopub.execute_input":"2023-01-14T08:00:38.387009Z","iopub.status.idle":"2023-01-14T08:00:49.395733Z","shell.execute_reply.started":"2023-01-14T08:00:38.38697Z","shell.execute_reply":"2023-01-14T08:00:49.394697Z"},"trusted":true},"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":"2023-01-14T08:00:49.397053Z","iopub.execute_input":"2023-01-14T08:00:49.397449Z","iopub.status.idle":"2023-01-14T08:00:49.469693Z","shell.execute_reply.started":"2023-01-14T08:00:49.397391Z","shell.execute_reply":"2023-01-14T08:00:49.468821Z"},"trusted":true},"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":"2023-01-14T08:00:49.471253Z","iopub.execute_input":"2023-01-14T08:00:49.471633Z","iopub.status.idle":"2023-01-14T08:05:13.848526Z","shell.execute_reply.started":"2023-01-14T08:00:49.471598Z","shell.execute_reply":"2023-01-14T08:05:13.847152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds","metadata":{"execution":{"iopub.status.busy":"2023-01-14T08:05:13.850549Z","iopub.execute_input":"2023-01-14T08:05:13.855633Z","iopub.status.idle":"2023-01-14T08:05:13.888342Z","shell.execute_reply.started":"2023-01-14T08:05:13.855539Z","shell.execute_reply":"2023-01-14T08:05:13.887482Z"},"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":[]}]}