{"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":"# Simple SegNet Baseline (Train)\n\nThis is the training part of the two part SegNet Baseline for this competition.\n\nYou can find the notebook to create the dataset used for training [here][2] to get a better understanding of how everything works.\n* Smp library is used to get the unet model.\n* EfficientNetB4 is used as the backbone initialized on imagenet weight.\n* Ash color images are used for training (With only the labeled frames and human_pixel_masks.\n* Custom implementation of dice score is used according to this competition.\n* After training, we find the best threshold for the valid set, which will then be used for the submission.\n* Wandb can also be used with this notebook to log experiments, just uncomment the wandb code snippets.\n\n### Please upvote if you find this useful.\n\n[1]: https://www.kaggle.com/code/shashwatraman/simple-unet-baseline-infer-lb-0-580\n[2]: https://www.kaggle.com/code/shashwatraman/contrails-dataset-ash-color/notebook\n","metadata":{}},{"cell_type":"markdown","source":"## Import Libraries","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nimport os\nimport random\nimport math\nfrom collections import defaultdict\nimport cv2\nimport skimage\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport torch\nfrom torch import nn\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nfrom transformers import get_cosine_schedule_with_warmup\n\ntorch.__version__","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-18T11:00:07.464939Z","iopub.execute_input":"2023-06-18T11:00:07.465758Z","iopub.status.idle":"2023-06-18T11:00:12.434872Z","shell.execute_reply.started":"2023-06-18T11:00:07.465723Z","shell.execute_reply":"2023-06-18T11:00:12.433948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nimport sys\nsys.path.append(\"../input/pretrained-models-pytorch\")\nsys.path.append(\"../input/efficientnet-pytorch\")\nsys.path.append(\"/kaggle/input/smp-github/segmentation_models.pytorch-master\")\nimport segmentation_models_pytorch as smp\n\nprint(f\"Segmentation Models version: {smp.__version__}\")","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip install -qU wandb\n#import wandb\n#wandb.login(key='')","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:14.470595Z","iopub.execute_input":"2023-06-18T11:00:14.471277Z","iopub.status.idle":"2023-06-18T11:00:14.475364Z","shell.execute_reply.started":"2023-06-18T11:00:14.471242Z","shell.execute_reply":"2023-06-18T11:00:14.474482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Preparation","metadata":{}},{"cell_type":"code","source":"class Config:\n    train = True\n    \n    num_epochs = 10\n    num_classes = 1\n    batch_size = 32//8\n    seed = 42\n    #net = \"deeplabv3\"\n    #net = 'unet'\n    #net = 'DeepLabV3Plus'\n    net = 'unetplusplus'\n    encoder = 'efficientnet-b4'\n    #encoder = 'resnet50'\n    #pretrained = True\n    weights = 'imagenet'\n    classes = ['contrail']\n    activation = None\n    in_chans = 3\n    \n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    image_size = 512\n    warmup = 0.01\n    lr = 3e-3\n    \n    const_lr = True\n    \nclass Paths:\n    data_root = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\n    contrails = '/kaggle/input/contrails-images-ash-color/contrails/'\n    train_path = '/kaggle/input/contrails-images-ash-color/train_df.csv'\n    valid_path = '/kaggle/input/contrails-images-ash-color/valid_df.csv'","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:14.478402Z","iopub.execute_input":"2023-06-18T11:00:14.478738Z","iopub.status.idle":"2023-06-18T11:00:14.51196Z","shell.execute_reply.started":"2023-06-18T11:00:14.478706Z","shell.execute_reply":"2023-06-18T11:00:14.51112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed=1234):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    \n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:14.513192Z","iopub.execute_input":"2023-06-18T11:00:14.513515Z","iopub.status.idle":"2023-06-18T11:00:14.530881Z","shell.execute_reply.started":"2023-06-18T11:00:14.513483Z","shell.execute_reply":"2023-06-18T11:00:14.529991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"set_seed(Config.seed)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:14.532211Z","iopub.execute_input":"2023-06-18T11:00:14.532544Z","iopub.status.idle":"2023-06-18T11:00:14.541868Z","shell.execute_reply.started":"2023-06-18T11:00:14.532514Z","shell.execute_reply":"2023-06-18T11:00:14.540983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Import dataframes\ntrain_df = pd.read_csv(Paths.train_path)\nvalid_df = pd.read_csv(Paths.valid_path)\n\ntrain_df['path'] = Paths.contrails + train_df['record_id'].astype(str) + '.npy'\nvalid_df['path'] = Paths.contrails + valid_df['record_id'].astype(str) + '.npy'\n\ntrain_df.shape, valid_df.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:14.543234Z","iopub.execute_input":"2023-06-18T11:00:14.543763Z","iopub.status.idle":"2023-06-18T11:00:14.590388Z","shell.execute_reply.started":"2023-06-18T11:00:14.543732Z","shell.execute_reply":"2023-06-18T11:00:14.589226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, train=True):\n        \n        self.df = df\n        self.trn = train\n    \n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        con_path = row.path\n        con = np.load(str(con_path))\n        \n        img = con[..., :-1]\n        label = con[..., -1]\n        \n        img = torch.tensor(img)\n        label = torch.tensor(label)\n        \n        img = img.permute(2, 0, 1)\n            \n        return img.float(), label.float()\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:14.591972Z","iopub.execute_input":"2023-06-18T11:00:14.592369Z","iopub.status.idle":"2023-06-18T11:00:14.6016Z","shell.execute_reply.started":"2023-06-18T11:00:14.592337Z","shell.execute_reply":"2023-06-18T11:00:14.600509Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_ds = ContrailsDataset(\n        train_df,\n        train=True\n    )\n\nvalid_ds = ContrailsDataset(\n        valid_df,\n        train=False\n    )\n\ntrain_dl = DataLoader(train_ds, batch_size=Config.batch_size , shuffle=True, num_workers = 2)    \nvalid_dl = DataLoader(valid_ds, batch_size=1, num_workers = 2)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:14.603074Z","iopub.execute_input":"2023-06-18T11:00:14.603553Z","iopub.status.idle":"2023-06-18T11:00:14.612971Z","shell.execute_reply.started":"2023-06-18T11:00:14.603517Z","shell.execute_reply":"2023-06-18T11:00:14.611591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, label = next(iter(train_dl))\nimg.shape, label.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:14.624071Z","iopub.execute_input":"2023-06-18T11:00:14.626459Z","iopub.status.idle":"2023-06-18T11:00:14.954871Z","shell.execute_reply.started":"2023-06-18T11:00:14.626426Z","shell.execute_reply":"2023-06-18T11:00:14.953541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img, label = next(iter(valid_dl))\nimg.shape, label.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:14.956413Z","iopub.execute_input":"2023-06-18T11:00:14.960082Z","iopub.status.idle":"2023-06-18T11:00:15.152502Z","shell.execute_reply.started":"2023-06-18T11:00:14.960044Z","shell.execute_reply":"2023-06-18T11:00:15.149942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def display_random_images(dataset, n=10, seed=None):\n    if seed:\n        random.seed(seed)\n    random_samples_idx = random.sample(range(len(dataset)), k=n)\n    plt.figure(figsize=(30, 20))\n    \n    for i, targ_sample in enumerate(random_samples_idx):\n        targ_image, targ_label = dataset[targ_sample][0], dataset[targ_sample][1]\n        \n        targ_image = targ_image.permute(1, 2, 0)\n        \n        plt.subplot(1, n, i+1)\n        plt.imshow(targ_image)\n        plt.axis(False)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:15.156996Z","iopub.execute_input":"2023-06-18T11:00:15.159247Z","iopub.status.idle":"2023-06-18T11:00:15.177939Z","shell.execute_reply.started":"2023-06-18T11:00:15.159208Z","shell.execute_reply":"2023-06-18T11:00:15.177005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_random_images(train_ds, 4, 42)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:15.180795Z","iopub.execute_input":"2023-06-18T11:00:15.181163Z","iopub.status.idle":"2023-06-18T11:00:16.597921Z","shell.execute_reply.started":"2023-06-18T11:00:15.181131Z","shell.execute_reply":"2023-06-18T11:00:16.596903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_random_images(valid_ds, 4, 42)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:16.599252Z","iopub.execute_input":"2023-06-18T11:00:16.599646Z","iopub.status.idle":"2023-06-18T11:00:17.850636Z","shell.execute_reply.started":"2023-06-18T11:00:16.599615Z","shell.execute_reply":"2023-06-18T11:00:17.849791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def dice_coef(y_true, y_pred, thr=0.5, epsilon=0.001):\n    y_true = y_true.flatten()\n    y_pred = (y_pred>thr).astype(np.float32).flatten()\n    inter = (y_true*y_pred).sum()\n    den = y_true.sum() + y_pred.sum()\n    dice = ((2*inter+epsilon)/(den+epsilon))\n    return dice","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.852099Z","iopub.execute_input":"2023-06-18T11:00:17.852686Z","iopub.status.idle":"2023-06-18T11:00:17.858808Z","shell.execute_reply.started":"2023-06-18T11:00:17.852654Z","shell.execute_reply":"2023-06-18T11:00:17.857679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fl_coef(y_true, y_pred, logist, reduce=True, alpha = 1.0, gamma=2.0):\n    y_true = y_true.flatten()\n    y_pred = y_pred.flatten()\n    bce = torch.nn.torch.nn.functional.binary_cross_entropy_with_logits(y_pred,y_true)\n    pt = torch.exp(-bce)\n    fl = alpha * (1-pt)**gamma * bce\n    if reduce:\n        fl = torch.mean(fl)\n    return fl","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.860312Z","iopub.execute_input":"2023-06-18T11:00:17.860964Z","iopub.status.idle":"2023-06-18T11:00:17.86959Z","shell.execute_reply.started":"2023-06-18T11:00:17.860901Z","shell.execute_reply":"2023-06-18T11:00:17.868774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SegNet(nn.Module):\n    def __init__(self, cfg):\n        super(SegNet, self).__init__()\n        \n        self.cfg = cfg\n        self.training = True\n        \n        if cfg.net.lower() == \"unet\":\n            self.model = smp.Unet(\n                encoder_name=cfg.encoder, \n                encoder_weights=cfg.weights, \n                decoder_use_batchnorm=True,\n                classes=len(cfg.classes), \n                activation=cfg.activation,\n            )\n            print('unet')\n        elif cfg.net.lower() == 'deeplabv3':\n            self.model = smp.DeepLabV3(\n                encoder_name =cfg.encoder,\n                encoder_weights=cfg.weights,    # use `imagenet` pre-trained weights for encoder initialization\n                in_channels=3,                  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n                classes=len(cfg.classes),        # model output channels (number of classes in your dataset)\n                activation=cfg.activation,\n                #decoder_use_batchnorm=True\n                )\n            print('deeplabv3')\n        elif cfg.net.lower() == \"deeplabv3plus\":\n            self.model = smp.DeepLabV3Plus(\n                encoder_name =cfg.encoder,\n                encoder_weights=cfg.weights,    # use `imagenet` pre-trained weights for encoder initialization\n                in_channels=3,                  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n                classes=len(cfg.classes),        # model output channels (number of classes in your dataset)\n                activation=cfg.activation,\n                #decoder_use_batchnorm=True\n                )\n            print(\"deeplabv3+\")\n        elif cfg.net.lower() == 'unetplusplus':\n            self.model = smp.UnetPlusPlus(\n                encoder_name =cfg.encoder,\n                encoder_weights=cfg.weights,    # use `imagenet` pre-trained weights for encoder initialization\n                in_channels=3,                  # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n                classes=len(cfg.classes),        # model output channels (number of classes in your dataset)\n                activation=cfg.activation,\n                decoder_use_batchnorm=True\n                )\n            print('unet++')\n        else:\n            raise Exception(\"error net name\")\n        \n        self.loss_fn = [smp.losses.DiceLoss(mode='binary'),smp.losses.FocalLoss(mode=\"binary\")]\n    \n    def forward(self, imgs, targets):\n        #print('imgs:',imgs.shape, ' targets:',targets.shape)\n        x = imgs\n        y = targets\n\n        if Config.image_size != 256:\n            x = torch.nn.functional.interpolate(x, \n                                                size=Config.image_size,\n                                                mode='bilinear'\n                                               )    \n            \n        logits = self.model(x)\n        if Config.image_size != 256:\n            logits = torch.nn.functional.interpolate(logits, \n                                                size=256,\n                                                mode='bilinear'\n                                               )  \n        loss = 0\n        for fn in self.loss_fn:\n            loss += fn(logits,y)\n            #print(loss.shape,fn)\n        loss /= len(self.loss_fn)\n        #loss = torch.mean([fn(logits,y) for fn in self.loss_fn])\n        \n        return {\"loss\": loss, \"logits\": logits.sigmoid(), \"logits_raw\": logits, \"target\": y}","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.87111Z","iopub.execute_input":"2023-06-18T11:00:17.871754Z","iopub.status.idle":"2023-06-18T11:00:17.890385Z","shell.execute_reply.started":"2023-06-18T11:00:17.871724Z","shell.execute_reply":"2023-06-18T11:00:17.889418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_step(model, dataloader, optimizer, scheduler, device):\n    \n    model.train()\n    \n    train_losses = []\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    \n    for step, (X, y) in pbar:\n        \n        X, y = X.to(device), y.to(device)\n        \n   \n        \n        torch.set_grad_enabled(True)\n        \n        output_dict = model(X, y)\n        loss = output_dict[\"loss\"]\n        train_losses.append(loss.item())\n        \n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n        \n        if scheduler is not None:\n            scheduler.step()\n    \n    train_loss = np.sum(train_losses)\n    \n    return train_loss","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.892092Z","iopub.execute_input":"2023-06-18T11:00:17.892689Z","iopub.status.idle":"2023-06-18T11:00:17.903671Z","shell.execute_reply.started":"2023-06-18T11:00:17.89266Z","shell.execute_reply":"2023-06-18T11:00:17.90278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_step(model, dataloader, device):\n    \n    model.eval()\n    torch.set_grad_enabled(False)\n    \n    val_data = defaultdict(list)\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid')\n    for step, (X, y) in pbar: \n        X, y = X.to(device), y.to(device)\n\n        output = model(X, y)\n        for key, val in output.items():\n            val_data[key] += [output[key]]\n\n    for key, val in output.items():\n        value = val_data[key]\n        if len(value[0].shape) == 0:\n            val_data[key] = torch.stack(value)\n        else:\n            val_data[key] = torch.cat(value, dim=0).cpu().detach().numpy()\n    \n    val_losses = val_data[\"loss\"].cpu().numpy()\n    val_loss = np.sum(val_losses)\n    \n    val_dice = dice_coef(val_data['target'], val_data['logits'])\n    \n    return val_loss, val_dice","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.905198Z","iopub.execute_input":"2023-06-18T11:00:17.905888Z","iopub.status.idle":"2023-06-18T11:00:17.916919Z","shell.execute_reply.started":"2023-06-18T11:00:17.905855Z","shell.execute_reply":"2023-06-18T11:00:17.915896Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.auto import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.918627Z","iopub.execute_input":"2023-06-18T11:00:17.919043Z","iopub.status.idle":"2023-06-18T11:00:17.925926Z","shell.execute_reply.started":"2023-06-18T11:00:17.919013Z","shell.execute_reply":"2023-06-18T11:00:17.924865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, train_dataloader, test_dataloader, optimizer,scheduler, epochs, device):\n    results = {'train_loss': [],\n              'val_loss': [],\n              'val_dice': []}\n    for epoch in range(epochs):\n        \n        set_seed(Config.seed + epoch)\n        print(\"EPOCH:\", epoch)\n        \n        train_loss = train_step(model,\n                              train_dataloader,\n                              optimizer,scheduler,\n                              device)\n        val_loss, val_dice = test_step(model,\n                            test_dataloader,\n                            device)\n        \n        train_loss = train_loss / len(train_ds)\n        val_loss = val_loss / len(valid_ds)\n        \n        print(f'Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | Val Dice: {val_dice:.4f}')\n        print(f\"Learning rate: {optimizer.param_groups[0]['lr']}\")\n        \n        results['train_loss'].append(train_loss)\n        results['val_loss'].append(val_loss)\n        results['val_dice'].append(val_dice)\n        \n        #wandb.log({\n        ##    \"Train Loss\": train_loss,\n        # #   \"Valid Loss\": val_loss,\n         #   'Valid Dice': val_dice})\n        \n        PATH = f\"epoch-{epoch}.pth\"\n        torch.save(model.state_dict(), PATH)\n        #wandb.save(PATH)\n\n    return results","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.927168Z","iopub.execute_input":"2023-06-18T11:00:17.928257Z","iopub.status.idle":"2023-06-18T11:00:17.938095Z","shell.execute_reply.started":"2023-06-18T11:00:17.928224Z","shell.execute_reply":"2023-06-18T11:00:17.937102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_optimizer(lr, params):\n    \n    model_optimizer = torch.optim.Adam(\n            filter(lambda p: p.requires_grad, params), \n            lr=lr,\n            weight_decay=0)\n    \n    return model_optimizer","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.93976Z","iopub.execute_input":"2023-06-18T11:00:17.940582Z","iopub.status.idle":"2023-06-18T11:00:17.949713Z","shell.execute_reply.started":"2023-06-18T11:00:17.940552Z","shell.execute_reply":"2023-06-18T11:00:17.948972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_scheduler(cfg, optimizer, total_steps):\n    scheduler = get_cosine_schedule_with_warmup(\n        optimizer,\n        num_warmup_steps= cfg.warmup * (total_steps // cfg.batch_size),\n        num_training_steps= cfg.num_epochs * (total_steps // cfg.batch_size)\n    )\n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.952965Z","iopub.execute_input":"2023-06-18T11:00:17.953263Z","iopub.status.idle":"2023-06-18T11:00:17.960272Z","shell.execute_reply.started":"2023-06-18T11:00:17.953235Z","shell.execute_reply":"2023-06-18T11:00:17.959598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_EPOCHS = Config.num_epochs\nmodel = SegNet(Config).to(Config.device)\n\n#run = wandb.init(project='Google Contrails', \n#                 config={k:v for k, v in dict(vars(Config)).items() if '__' not in k},\n#                 name=f\"{Config.encoder}-{Config.num_epochs}epos-{Config.lr}-segnet\"\n#                )\n\ntotal_steps = len(train_ds)\noptimizer = get_optimizer(lr=Config.lr, params=model.parameters())\nif Config.const_lr:\n    scheduler = None\nelse:\n    scheduler = get_scheduler(Config, optimizer, total_steps)\n\n#wandb.watch(model, log_freq=100, log='all')\n\nfrom timeit import default_timer as timer\nstart_time = timer()\n\nmodel_results = train(model, train_dl, valid_dl, optimizer,scheduler, NUM_EPOCHS, Config.device)\n#optimizer = get_optimizer(lr=0.00005, params=model.parameters())\n#model_results = train(model, train_dl, valid_dl, optimizer, 5, Config.device)\n\nend_time = timer()\n\n#run.finish()\nprint(f'Total Training Time: {end_time-start_time:.3f} seconds')","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:17.961328Z","iopub.execute_input":"2023-06-18T11:00:17.962157Z","iopub.status.idle":"2023-06-18T11:00:42.212456Z","shell.execute_reply.started":"2023-06-18T11:00:17.962126Z","shell.execute_reply":"2023-06-18T11:00:42.21079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Finding the Best Threshold","metadata":{}},{"cell_type":"code","source":"# Predicting the Valid Set\nmodel.eval()\ntorch.set_grad_enabled(False)\n\nval_data = defaultdict(list)\npbar = tqdm(enumerate(valid_dl), total=len(valid_dl), desc='Valid')\nfor step, (X, y) in pbar: \n    X, y = X.to(Config.device), y.to(Config.device)\n    #print(X.shape, y.shape)\n    output = model(X, y)\n    #print(len(output))\n    for key, val in output.items():\n        val_data[key] += [output[key]]\n\nfor key, val in output.items():\n    value = val_data[key]\n    if len(value[0].shape) == 0:\n        val_data[key] = torch.stack(value)\n    else:\n        val_data[key] = torch.cat(value, dim=0).cpu().detach().numpy()\n\nval_losses = val_data[\"loss\"].cpu().numpy()\nval_loss = np.sum(val_losses)\nval_loss = val_loss / len(valid_ds)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:42.213952Z","iopub.status.idle":"2023-06-18T11:00:42.214658Z","shell.execute_reply.started":"2023-06-18T11:00:42.214416Z","shell.execute_reply":"2023-06-18T11:00:42.214443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = val_data['logits']\nground_truths = val_data['target']","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:42.215965Z","iopub.status.idle":"2023-06-18T11:00:42.216659Z","shell.execute_reply.started":"2023-06-18T11:00:42.216422Z","shell.execute_reply":"2023-06-18T11:00:42.216445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions.shape, ground_truths.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:42.217904Z","iopub.status.idle":"2023-06-18T11:00:42.218599Z","shell.execute_reply.started":"2023-06-18T11:00:42.218346Z","shell.execute_reply":"2023-06-18T11:00:42.218378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Finding the Best Threshold\nbdice = -1\nbi = None\nfor i in tqdm(np.arange(0, 1.01, 0.01)):\n    val_dice = dice_coef(ground_truths, predictions, i)\n    if val_dice > bdice:\n        bdice = val_dice\n        bi = i","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:42.219854Z","iopub.status.idle":"2023-06-18T11:00:42.220543Z","shell.execute_reply.started":"2023-06-18T11:00:42.220293Z","shell.execute_reply":"2023-06-18T11:00:42.220316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'Best Threshold: {bi}')\nprint(f'Best Validation Dice Score: {bdice}')","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:42.221761Z","iopub.status.idle":"2023-06-18T11:00:42.222477Z","shell.execute_reply.started":"2023-06-18T11:00:42.222226Z","shell.execute_reply":"2023-06-18T11:00:42.222249Z"},"trusted":true},"execution_count":null,"outputs":[]}]}