{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":51753,"databundleVersionId":5692552,"sourceType":"competition"},{"sourceId":5848162,"sourceType":"datasetVersion","datasetId":3362727}],"dockerImageVersionId":30498,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Simple Unet Baseline (Train)\n\nThis is the training part of the two part Unet Baseline for this competition.\n#### Inference Notebook: [Simple Unet Baseline (Infer)][1]. \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* EfficientNetB0 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**Version 5** Updates:\n* Added some Augmentations\n* Trained for more Epochs\n* Option to increase image size\n\n### Please upvote if you find this useful.\n\n[1]: https://www.kaggle.com/code/shashwatraman/simple-unet-pytorch-baseline-infer\n[2]: https://www.kaggle.com/code/shashwatraman/contrails-dataset-ash-color/notebook","metadata":{}},{"cell_type":"markdown","source":"## Import Libraries","metadata":{}},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\nimport segmentation_models_pytorch as smp","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2025-02-12T10:06:11.806511Z","iopub.execute_input":"2025-02-12T10:06:11.8072Z","iopub.status.idle":"2025-02-12T10:06:30.323558Z","shell.execute_reply.started":"2025-02-12T10:06:11.807172Z","shell.execute_reply":"2025-02-12T10:06:30.322774Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !pip install -qU wandb\n# import wandb\n# wandb.login(key='')","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:30.325377Z","iopub.execute_input":"2025-02-12T10:06:30.326007Z","iopub.status.idle":"2025-02-12T10:06:30.329941Z","shell.execute_reply.started":"2025-02-12T10:06:30.325975Z","shell.execute_reply":"2025-02-12T10:06:30.328952Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pathlib opencv-python-headless scikit-image numpy pandas matplotlib torch torchvision albumentations pillow tqdm transformers","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T10:06:30.331016Z","iopub.execute_input":"2025-02-12T10:06:30.331537Z","iopub.status.idle":"2025-02-12T10:06:38.260333Z","shell.execute_reply.started":"2025-02-12T10:06:30.331512Z","shell.execute_reply":"2025-02-12T10:06:38.25947Z"}},"outputs":[],"execution_count":null},{"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\nimport albumentations as A\nimport torch.nn.functional as F\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":"2025-02-12T10:06:38.262657Z","iopub.execute_input":"2025-02-12T10:06:38.262941Z","iopub.status.idle":"2025-02-12T10:06:48.373723Z","shell.execute_reply.started":"2025-02-12T10:06:38.262916Z","shell.execute_reply":"2025-02-12T10:06:48.372871Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Preparation","metadata":{}},{"cell_type":"code","source":"class Config:\n    train = True\n    train_aug=True\n    \n    num_epochs = 30\n    num_classes = 1\n    batch_size = 32\n    seed = 42\n    \n    encoder = 'efficientnet-b3'\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 = 256\n    warmup = 0\n    lr = 3e-3\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":"2025-02-12T10:06:48.37478Z","iopub.execute_input":"2025-02-12T10:06:48.375531Z","iopub.status.idle":"2025-02-12T10:06:48.4058Z","shell.execute_reply.started":"2025-02-12T10:06:48.375497Z","shell.execute_reply":"2025-02-12T10:06:48.404956Z"},"trusted":true},"outputs":[],"execution_count":null},{"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 = True\n    torch.backends.cudnn.benchmark = False","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:48.407009Z","iopub.execute_input":"2025-02-12T10:06:48.407263Z","iopub.status.idle":"2025-02-12T10:06:48.422646Z","shell.execute_reply.started":"2025-02-12T10:06:48.407241Z","shell.execute_reply":"2025-02-12T10:06:48.421976Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"set_seed(9)","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:48.423771Z","iopub.execute_input":"2025-02-12T10:06:48.42402Z","iopub.status.idle":"2025-02-12T10:06:48.437138Z","shell.execute_reply.started":"2025-02-12T10:06:48.423999Z","shell.execute_reply":"2025-02-12T10:06:48.436459Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":"2025-02-12T10:06:48.438099Z","iopub.execute_input":"2025-02-12T10:06:48.438302Z","iopub.status.idle":"2025-02-12T10:06:48.49984Z","shell.execute_reply.started":"2025-02-12T10:06:48.438284Z","shell.execute_reply":"2025-02-12T10:06:48.49898Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transform_size = A.Compose([\n    A.Resize(Config.image_size, Config.image_size, interpolation=cv2.INTER_LANCZOS4, always_apply=True)\n])\n\ntrain_transform = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.RandomResizedCrop(height=256, width=256, scale=(0.75, 1.0), p=0.6)\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T10:06:48.500748Z","iopub.execute_input":"2025-02-12T10:06:48.500975Z","iopub.status.idle":"2025-02-12T10:06:48.506053Z","shell.execute_reply.started":"2025-02-12T10:06:48.500956Z","shell.execute_reply":"2025-02-12T10:06:48.505191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, train=True, transform=None):\n        \n        self.df = df\n        self.trn = train\n        self.transform = transform\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 = img.astype(np.float32)\n        label = label.astype(np.float32)\n        \n        if Config.train_aug:\n            if self.transform is not None:\n                augmented = self.transform(image=img, mask=label)\n                img = augmented['image']\n                label = augmented['mask']\n                \n        if Config.image_size != 256:\n            img = transform_size(image=img)[\"image\"]\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":"2025-02-12T10:06:48.508774Z","iopub.execute_input":"2025-02-12T10:06:48.508993Z","iopub.status.idle":"2025-02-12T10:06:48.52117Z","shell.execute_reply.started":"2025-02-12T10:06:48.508975Z","shell.execute_reply":"2025-02-12T10:06:48.520407Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = ContrailsDataset(\n        train_df,\n        train=True,\n        transform=train_transform\n    )\n\nvalid_ds = ContrailsDataset(\n        valid_df,\n        train=False,\n        transform=None\n    )\n\ntrain_dl = DataLoader(train_ds, batch_size=Config.batch_size , shuffle=True, num_workers = 2)    \nvalid_dl = DataLoader(valid_ds, batch_size=Config.batch_size, num_workers = 2)","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:48.522027Z","iopub.execute_input":"2025-02-12T10:06:48.522284Z","iopub.status.idle":"2025-02-12T10:06:48.536997Z","shell.execute_reply.started":"2025-02-12T10:06:48.522264Z","shell.execute_reply":"2025-02-12T10:06:48.536198Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img, label = next(iter(train_dl))\nimg.shape, label.shape","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:48.537871Z","iopub.execute_input":"2025-02-12T10:06:48.538135Z","iopub.status.idle":"2025-02-12T10:06:50.103227Z","shell.execute_reply.started":"2025-02-12T10:06:48.538114Z","shell.execute_reply":"2025-02-12T10:06:50.102318Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img, label = next(iter(valid_dl))\nimg.shape, label.shape","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:50.104537Z","iopub.execute_input":"2025-02-12T10:06:50.104829Z","iopub.status.idle":"2025-02-12T10:06:51.623148Z","shell.execute_reply.started":"2025-02-12T10:06:50.104801Z","shell.execute_reply":"2025-02-12T10:06:51.622117Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":"2025-02-12T10:06:51.624403Z","iopub.execute_input":"2025-02-12T10:06:51.624727Z","iopub.status.idle":"2025-02-12T10:06:51.630431Z","shell.execute_reply.started":"2025-02-12T10:06:51.624698Z","shell.execute_reply":"2025-02-12T10:06:51.629563Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"display_random_images(train_ds, 4, 42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T10:06:51.631383Z","iopub.execute_input":"2025-02-12T10:06:51.631647Z","iopub.status.idle":"2025-02-12T10:06:53.197736Z","shell.execute_reply.started":"2025-02-12T10:06:51.631623Z","shell.execute_reply":"2025-02-12T10:06:53.196793Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"display_random_images(valid_ds, 4, 42)","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:53.198941Z","iopub.execute_input":"2025-02-12T10:06:53.199204Z","iopub.status.idle":"2025-02-12T10:06:54.351179Z","shell.execute_reply.started":"2025-02-12T10:06:53.199182Z","shell.execute_reply":"2025-02-12T10:06:54.349948Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":"2025-02-12T10:06:54.352512Z","iopub.execute_input":"2025-02-12T10:06:54.353085Z","iopub.status.idle":"2025-02-12T10:06:54.358631Z","shell.execute_reply.started":"2025-02-12T10:06:54.353055Z","shell.execute_reply":"2025-02-12T10:06:54.3578Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, cfg):\n        super(UNet, self).__init__()\n        \n        self.cfg = cfg\n        self.training = True\n        \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        \n        self.loss_fn = smp.losses.DiceLoss(mode='binary')\n    \n    def forward(self, imgs, targets):\n        \n        x = imgs\n        y = targets\n\n        logits = self.model(x)\n        \n        if Config.image_size != 256:\n            logits = F.interpolate(logits, size=(256, 256), mode='nearest-exact')\n        \n        loss = self.loss_fn(logits, y)\n        \n        return {\"loss\": loss, \"logits\": logits.sigmoid(), \"logits_raw\": logits, \"target\": y}","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:54.35963Z","iopub.execute_input":"2025-02-12T10:06:54.359918Z","iopub.status.idle":"2025-02-12T10:06:54.369495Z","shell.execute_reply.started":"2025-02-12T10:06:54.359898Z","shell.execute_reply":"2025-02-12T10:06:54.368658Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_step(model, dataloader, optimizer, 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        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":"2025-02-12T10:06:54.370604Z","iopub.execute_input":"2025-02-12T10:06:54.370845Z","iopub.status.idle":"2025-02-12T10:06:54.379898Z","shell.execute_reply.started":"2025-02-12T10:06:54.370824Z","shell.execute_reply":"2025-02-12T10:06:54.379183Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":"2025-02-12T10:06:54.380809Z","iopub.execute_input":"2025-02-12T10:06:54.381055Z","iopub.status.idle":"2025-02-12T10:06:54.394226Z","shell.execute_reply.started":"2025-02-12T10:06:54.381036Z","shell.execute_reply":"2025-02-12T10:06:54.393438Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm.auto import tqdm","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:54.395247Z","iopub.execute_input":"2025-02-12T10:06:54.395564Z","iopub.status.idle":"2025-02-12T10:06:54.404276Z","shell.execute_reply.started":"2025-02-12T10:06:54.395536Z","shell.execute_reply":"2025-02-12T10:06:54.403453Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(model, train_dataloader, test_dataloader, optimizer, 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,\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        \n#         wandb.save(PATH)\n\n    return results","metadata":{"execution":{"iopub.status.busy":"2025-02-12T10:06:54.40549Z","iopub.execute_input":"2025-02-12T10:06:54.406405Z","iopub.status.idle":"2025-02-12T10:06:54.414534Z","shell.execute_reply.started":"2025-02-12T10:06:54.406374Z","shell.execute_reply":"2025-02-12T10:06:54.413773Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":"2025-02-12T10:06:54.415469Z","iopub.execute_input":"2025-02-12T10:06:54.415667Z","iopub.status.idle":"2025-02-12T10:06:54.428923Z","shell.execute_reply.started":"2025-02-12T10:06:54.41565Z","shell.execute_reply":"2025-02-12T10:06:54.428309Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":"2025-02-12T10:06:54.429977Z","iopub.execute_input":"2025-02-12T10:06:54.430207Z","iopub.status.idle":"2025-02-12T10:06:54.439388Z","shell.execute_reply.started":"2025-02-12T10:06:54.430188Z","shell.execute_reply":"2025-02-12T10:06:54.438614Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_EPOCHS = Config.num_epochs\nmodel = UNet(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}-unet\"\n#                     )\n\ntotal_steps = len(train_ds)\noptimizer = get_optimizer(lr=Config.lr, params=model.parameters())\nscheduler = 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, NUM_EPOCHS, 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":"2025-02-12T10:06:54.440294Z","iopub.execute_input":"2025-02-12T10:06:54.440543Z","iopub.status.idle":"2025-02-12T10:41:22.725486Z","shell.execute_reply.started":"2025-02-12T10:06:54.440523Z","shell.execute_reply":"2025-02-12T10:41:22.723869Z"},"trusted":true},"outputs":[],"execution_count":null},{"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\n    output = model(X, y)\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T10:41:22.726436Z","iopub.status.idle":"2025-02-12T10:41:22.72672Z","shell.execute_reply.started":"2025-02-12T10:41:22.726588Z","shell.execute_reply":"2025-02-12T10:41:22.726603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = val_data['logits']\nground_truths = val_data['target']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T10:41:22.727973Z","iopub.status.idle":"2025-02-12T10:41:22.728266Z","shell.execute_reply.started":"2025-02-12T10:41:22.728129Z","shell.execute_reply":"2025-02-12T10:41:22.728143Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions.shape, ground_truths.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T10:41:22.729729Z","iopub.status.idle":"2025-02-12T10:41:22.729986Z","shell.execute_reply.started":"2025-02-12T10:41:22.729863Z","shell.execute_reply":"2025-02-12T10:41:22.729875Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T12:41:01.682028Z","iopub.execute_input":"2025-02-12T12:41:01.682366Z","iopub.status.idle":"2025-02-12T12:41:01.720486Z","shell.execute_reply.started":"2025-02-12T12:41:01.682337Z","shell.execute_reply":"2025-02-12T12:41:01.719378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f'Best Threshold: {bi}')\nprint(f'Best Validation Dice Score: {bdice}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T10:41:22.732218Z","iopub.status.idle":"2025-02-12T10:41:22.732536Z","shell.execute_reply.started":"2025-02-12T10:41:22.732365Z","shell.execute_reply":"2025-02-12T10:41:22.732378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T11:42:27.468507Z","iopub.execute_input":"2025-02-12T11:42:27.469174Z","iopub.status.idle":"2025-02-12T11:42:40.464659Z","shell.execute_reply.started":"2025-02-12T11:42:27.469142Z","shell.execute_reply":"2025-02-12T11:42:40.46378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dice_coef(y_true, y_pred, thr=0.5, epsilon=0.001):\n    \"\"\"\n    Compute the Dice coefficient between y_true and y_pred,\n    optionally thresholding y_pred first.\n    \"\"\"\n    import numpy as np\n    \n    y_true = y_true.flatten()\n    if thr is not None:\n        y_pred = (y_pred > thr).astype(np.float32).flatten()\n    else:\n        y_pred = y_pred.flatten().astype(np.float32)\n\n    inter = (y_true * y_pred).sum()\n    den = y_true.sum() + y_pred.sum()\n    dice = (2.0 * inter + epsilon) / (den + epsilon)\n    return dice\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T12:56:04.585497Z","iopub.execute_input":"2025-02-12T12:56:04.585834Z","iopub.status.idle":"2025-02-12T12:56:04.591321Z","shell.execute_reply.started":"2025-02-12T12:56:04.585805Z","shell.execute_reply":"2025-02-12T12:56:04.590467Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Block [22]\nimport segmentation_models_pytorch as smp\nimport torch.nn.functional as F\nimport torch.nn as nn\n\nclass UNet(nn.Module):\n    def __init__(self, cfg):\n        super(UNet, self).__init__()\n        self.cfg = cfg\n        self.training = True\n        \n        # Create the UNet model from SMP\n        self.model = smp.Unet(\n            encoder_name=cfg.encoder,\n            encoder_weights=cfg.weights,\n            decoder_use_batchnorm=True,\n            classes=1,  # <-- changed from len(cfg.classes) to 1\n            activation=None,  # we handle sigmoid ourselves\n        )\n        \n        # Combine Dice + Soft BCE for more stable training\n        self.dice_loss_fn = smp.losses.DiceLoss(mode='binary', from_logits=True)\n        self.bce_loss_fn  = smp.losses.SoftBCEWithLogitsLoss()\n        \n    def forward(self, imgs, targets):\n        x = imgs\n        y = targets\n        \n        logits = self.model(x)\n\n        if self.cfg.image_size != 256:\n            logits = F.interpolate(\n                logits,\n                size=(256, 256),\n                mode='nearest-exact'\n            )\n            y = F.interpolate(\n                y,\n                size=(256, 256),\n                mode='nearest-exact'\n            )\n\n        # Combine dice + BCE\n        dice_loss = self.dice_loss_fn(logits, y)\n        bce_loss = self.bce_loss_fn(logits, y)\n        loss = dice_loss + bce_loss\n\n        return {\n            \"loss\": loss,\n            \"logits\": torch.sigmoid(logits),\n            \"logits_raw\": logits,\n            \"target\": y\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T12:59:20.821161Z","iopub.execute_input":"2025-02-12T12:59:20.822059Z","iopub.status.idle":"2025-02-12T12:59:20.829275Z","shell.execute_reply.started":"2025-02-12T12:59:20.822026Z","shell.execute_reply":"2025-02-12T12:59:20.828388Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np\nimport torch\n\n# For mixed-precision:\nfrom torch.cuda.amp import autocast, GradScaler\nscaler = GradScaler()\n\ndef train_step(model, dataloader, optimizer, device, scheduler=None):\n    model.train()\n    train_losses = []\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    \n    for step, (X, y) in pbar:\n        X, y = X.to(device), y.to(device)\n        \n        # Clear gradients\n        optimizer.zero_grad(set_to_none=True)\n        \n        with torch.set_grad_enabled(True):\n            # Use autocast for mixed precision\n            with autocast():\n                output_dict = model(X, y)\n                loss = output_dict[\"loss\"]\n\n            # Scale the loss for backprop\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n        train_losses.append(loss.item())\n\n        # Step the scheduler if given\n        if scheduler is not None:\n            scheduler.step()\n\n    train_loss = np.sum(train_losses)\n    return train_loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T12:59:26.202879Z","iopub.execute_input":"2025-02-12T12:59:26.203537Z","iopub.status.idle":"2025-02-12T12:59:26.209725Z","shell.execute_reply.started":"2025-02-12T12:59:26.203504Z","shell.execute_reply":"2025-02-12T12:59:26.20895Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Block [24]\nfrom collections import defaultdict\nimport torch\n\n@torch.no_grad()\ndef test_step(model, dataloader, device):\n    model.eval()\n    val_data = defaultdict(list)\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid')\n    \n    for step, (X, y) in pbar:\n        X, y = X.to(device), y.to(device)\n        \n        output = model(X, y)\n        \n        # Expand loss to shape [1] so we can safely cat\n        val_data[\"loss\"].append(output[\"loss\"].unsqueeze(0))\n        val_data[\"logits\"].append(output[\"logits\"])\n        val_data[\"logits_raw\"].append(output[\"logits_raw\"])\n        val_data[\"target\"].append(output[\"target\"])\n\n    # Concatenate all items on CPU\n    for key, val_list in val_data.items():\n        val_data[key] = torch.cat([v.cpu() for v in val_list], dim=0)\n\n    # Now 'loss' is a 1-D tensor instead of zero-dim, so cat() works\n    val_losses = val_data[\"loss\"].numpy()\n    val_loss = np.sum(val_losses)\n\n    # Evaluate dice score\n    val_dice = dice_coef(val_data['target'].numpy(), val_data['logits'].numpy())\n\n    return val_loss, val_dice\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T13:12:17.525981Z","iopub.execute_input":"2025-02-12T13:12:17.52674Z","iopub.status.idle":"2025-02-12T13:12:17.534298Z","shell.execute_reply.started":"2025-02-12T13:12:17.526707Z","shell.execute_reply":"2025-02-12T13:12:17.533381Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train(model, train_dataloader, test_dataloader, optimizer, epochs, device, scheduler=None):\n    results = {\n        'train_loss': [],\n        'val_loss': [],\n        'val_dice': []\n    }\n    \n    best_dice = -1.0  # for tracking best model\n    \n    for epoch in range(epochs):\n        # Optionally reseed each epoch for reproducibility\n        set_seed(Config.seed + epoch)\n        \n        print(\"EPOCH:\", epoch)\n        \n        # Train for one epoch\n        train_loss = train_step(\n            model,\n            train_dataloader,\n            optimizer,\n            device,\n            scheduler\n        )\n        \n        # Validate\n        val_loss, val_dice = test_step(\n            model,\n            test_dataloader,\n            device\n        )\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        # Record metrics\n        results['train_loss'].append(train_loss)\n        results['val_loss'].append(val_loss)\n        results['val_dice'].append(val_dice)\n        \n        # Save best model (maximize Dice)\n        if val_dice > best_dice:\n            best_dice = val_dice\n            PATH = f\"best_model_epoch-{epoch}.pth\"\n            torch.save(model.state_dict(), PATH)\n            print(f\"  --> Saved best model so far, dice = {best_dice:.4f}\")\n\n        # You could also log metrics to W&B here if desired\n        # wandb.log({\n        #    \"Train Loss\": train_loss,\n        #    \"Valid Loss\": val_loss,\n        #    \"Valid Dice\": val_dice\n        # })\n\n    return results\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T12:59:32.791955Z","iopub.execute_input":"2025-02-12T12:59:32.792753Z","iopub.status.idle":"2025-02-12T12:59:32.799403Z","shell.execute_reply.started":"2025-02-12T12:59:32.79272Z","shell.execute_reply":"2025-02-12T12:59:32.798571Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\ndef get_optimizer(lr, params):\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\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T12:59:37.388698Z","iopub.execute_input":"2025-02-12T12:59:37.389332Z","iopub.status.idle":"2025-02-12T12:59:37.393472Z","shell.execute_reply.started":"2025-02-12T12:59:37.3893Z","shell.execute_reply":"2025-02-12T12:59:37.392611Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### Block [28]\nfrom transformers import get_cosine_schedule_with_warmup\n\ndef get_scheduler(cfg, optimizer, total_steps):\n    scheduler = get_cosine_schedule_with_warmup(\n        optimizer,\n        num_warmup_steps=0,  # <-- replaced cfg.warmup * (...)\n        num_training_steps=cfg.num_epochs * (total_steps // cfg.batch_size)\n    )\n    return scheduler\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T13:00:47.206991Z","iopub.execute_input":"2025-02-12T13:00:47.20761Z","iopub.status.idle":"2025-02-12T13:00:47.211866Z","shell.execute_reply.started":"2025-02-12T13:00:47.207582Z","shell.execute_reply":"2025-02-12T13:00:47.210882Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_EPOCHS = Config.num_epochs\nmodel = UNet(Config).to(Config.device)\n\ntotal_steps = len(train_ds)\noptimizer = get_optimizer(lr=Config.lr, params=model.parameters())\nscheduler = get_scheduler(Config, optimizer, total_steps)\n\nstart_time = timer()\nmodel_results = train(\n    model,\n    train_dl,\n    valid_dl,\n    optimizer,\n    NUM_EPOCHS,\n    Config.device,\n    scheduler\n)\nend_time = timer()\n\nprint(f\"Total Training Time: {end_time - start_time:.3f} seconds\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T13:12:22.951283Z","iopub.execute_input":"2025-02-12T13:12:22.951606Z","iopub.status.idle":"2025-02-12T15:48:50.638978Z","shell.execute_reply.started":"2025-02-12T13:12:22.95158Z","shell.execute_reply":"2025-02-12T15:48:50.637784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.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\n    output = model(X, y)\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T16:08:49.45206Z","iopub.execute_input":"2025-02-12T16:08:49.452402Z","iopub.status.idle":"2025-02-12T16:08:57.956274Z","shell.execute_reply.started":"2025-02-12T16:08:49.452367Z","shell.execute_reply":"2025-02-12T16:08:57.95515Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = val_data['logits']\nground_truths = val_data['target']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T16:09:13.401143Z","iopub.execute_input":"2025-02-12T16:09:13.402006Z","iopub.status.idle":"2025-02-12T16:09:13.406155Z","shell.execute_reply.started":"2025-02-12T16:09:13.401968Z","shell.execute_reply":"2025-02-12T16:09:13.405309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions.shape, ground_truths.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T16:09:23.687906Z","iopub.execute_input":"2025-02-12T16:09:23.68826Z","iopub.status.idle":"2025-02-12T16:09:23.694906Z","shell.execute_reply.started":"2025-02-12T16:09:23.688234Z","shell.execute_reply":"2025-02-12T16:09:23.694131Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T16:09:33.300438Z","iopub.execute_input":"2025-02-12T16:09:33.301231Z","iopub.status.idle":"2025-02-12T16:10:57.630162Z","shell.execute_reply.started":"2025-02-12T16:09:33.301197Z","shell.execute_reply":"2025-02-12T16:10:57.629299Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f'Best Threshold: {bi}')\nprint(f'Best Validation Dice Score: {bdice}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-12T16:11:34.164893Z","iopub.execute_input":"2025-02-12T16:11:34.165204Z","iopub.status.idle":"2025-02-12T16:11:34.169682Z","shell.execute_reply.started":"2025-02-12T16:11:34.16518Z","shell.execute_reply":"2025-02-12T16:11:34.168899Z"}},"outputs":[],"execution_count":null}]}