{"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":"\"\"\"\nLovasz-Softmax and Jaccard hinge loss in PyTorch\nMaxim Berman 2018 ESAT-PSI KU Leuven (MIT License)\n\"\"\"\n\nimport torch\nfrom torch.autograd import Variable\nimport torch.nn.functional as F\nimport numpy as np\ntry:\n    from itertools import  ifilterfalse\nexcept ImportError: # py3k\n    from itertools import  filterfalse\n\n\ndef lovasz_grad(gt_sorted):\n    \"\"\"\n    Computes gradient of the Lovasz extension w.r.t sorted errors\n    See Alg. 1 in paper\n    \"\"\"\n    p = len(gt_sorted)\n    gts = gt_sorted.sum()\n    intersection = gts - gt_sorted.float().cumsum(0)\n    union = gts + (1 - gt_sorted).float().cumsum(0)\n    jaccard = 1. - intersection / union\n    if p > 1: # cover 1-pixel case\n        jaccard[1:p] = jaccard[1:p] - jaccard[0:-1]\n    return jaccard\n\n\ndef iou_binary(preds, labels, EMPTY=1., ignore=None, per_image=True):\n    \"\"\"\n    IoU for foreground class\n    binary: 1 foreground, 0 background\n    \"\"\"\n    if not per_image:\n        preds, labels = (preds,), (labels,)\n    ious = []\n    for pred, label in zip(preds, labels):\n        intersection = ((label == 1) & (pred == 1)).sum()\n        union = ((label == 1) | ((pred == 1) & (label != ignore))).sum()\n        if not union:\n            iou = EMPTY\n        else:\n            iou = float(intersection) / union\n        ious.append(iou)\n    iou = f_mean(ious)    # mean accross images if per_image\n    return 100 * iou\n\n\ndef iou(preds, labels, C, EMPTY=1., ignore=None, per_image=False):\n    \"\"\"\n    Array of IoU for each (non ignored) class\n    \"\"\"\n    if not per_image:\n        preds, labels = (preds,), (labels,)\n    ious = []\n    for pred, label in zip(preds, labels):\n        iou = []    \n        for i in range(C):\n            if i != ignore: # The ignored label is sometimes among predicted classes (ENet - CityScapes)\n                intersection = ((label == i) & (pred == i)).sum()\n                union = ((label == i) | ((pred == i) & (label != ignore))).sum()\n                if not union:\n                    iou.append(EMPTY)\n                else:\n                    iou.append(float(intersection) / union)\n        ious.append(iou)\n    ious = map(f_mean, zip(*ious)) # mean accross images if per_image\n    return 100 * np.array(ious)\n\n\n# --------------------------- BINARY LOSSES ---------------------------\n\n\ndef lovasz_hinge(logits, labels, per_image=True, ignore=None):\n    \"\"\"\n    Binary Lovasz hinge loss\n      logits: [B, H, W] Variable, logits at each pixel (between -\\infty and +\\infty)\n      labels: [B, H, W] Tensor, binary ground truth masks (0 or 1)\n      per_image: compute the loss per image instead of per batch\n      ignore: void class id\n    \"\"\"\n    if per_image:\n        loss = f_mean(lovasz_hinge_flat(*flatten_binary_scores(log.unsqueeze(0), lab.unsqueeze(0), ignore))\n                          for log, lab in zip(logits, labels))\n    else:\n        loss = lovasz_hinge_flat(*flatten_binary_scores(logits, labels, ignore))\n    return loss\n\n\ndef lovasz_hinge_flat(logits, labels):\n    \"\"\"\n    Binary Lovasz hinge loss\n      logits: [P] Variable, logits at each prediction (between -\\infty and +\\infty)\n      labels: [P] Tensor, binary ground truth labels (0 or 1)\n      ignore: label to ignore\n    \"\"\"\n    if len(labels) == 0:\n        # only void pixels, the gradients should be 0\n        return logits.sum() * 0.\n    signs = 2. * labels.float() - 1.\n    errors = (1. - logits * Variable(signs))\n    errors_sorted, perm = torch.sort(errors, dim=0, descending=True)\n    perm = perm.data\n    gt_sorted = labels[perm]\n    grad = lovasz_grad(gt_sorted)\n    #loss = torch.dot(F.relu(errors_sorted), Variable(grad))\n    loss = torch.dot(F.elu(errors_sorted)+1, Variable(grad))\n    return loss\n\n\ndef flatten_binary_scores(scores, labels, ignore=None):\n    \"\"\"\n    Flattens predictions in the batch (binary case)\n    Remove labels equal to 'ignore'\n    \"\"\"\n    scores = scores.view(-1)\n    labels = labels.view(-1)\n    if ignore is None:\n        return scores, labels\n    valid = (labels != ignore)\n    vscores = scores[valid]\n    vlabels = labels[valid]\n    return vscores, vlabels\n\n\nclass StableBCELoss(torch.nn.modules.Module):\n    def __init__(self):\n         super(StableBCELoss, self).__init__()\n    def forward(self, input, target):\n         neg_abs = - input.abs()\n         loss = input.clamp(min=0) - input * target + (1 + neg_abs.exp()).log()\n         return loss.mean()\n\n\ndef binary_xloss(logits, labels, ignore=None):\n    \"\"\"\n    Binary Cross entropy loss\n      logits: [B, H, W] Variable, logits at each pixel (between -\\infty and +\\infty)\n      labels: [B, H, W] Tensor, binary ground truth masks (0 or 1)\n      ignore: void class id\n    \"\"\"\n    logits, labels = flatten_binary_scores(logits, labels, ignore)\n    loss = StableBCELoss()(logits, Variable(labels.float()))\n    return loss\n\n\n# --------------------------- MULTICLASS LOSSES ---------------------------\n\n\ndef lovasz_softmax(probas, labels, only_present=False, per_image=False, ignore=None):\n    \"\"\"\n    Multi-class Lovasz-Softmax loss\n      probas: [B, C, H, W] Variable, class probabilities at each prediction (between 0 and 1)\n      labels: [B, H, W] Tensor, ground truth labels (between 0 and C - 1)\n      only_present: average only on classes present in ground truth\n      per_image: compute the loss per image instead of per batch\n      ignore: void class labels\n    \"\"\"\n    if per_image:\n        loss = f_mean(lovasz_softmax_flat(*flatten_probas(prob.unsqueeze(0), lab.unsqueeze(0), ignore), only_present=only_present)\n                          for prob, lab in zip(probas, labels))\n    else:\n        loss = lovasz_softmax_flat(*flatten_probas(probas, labels, ignore), only_present=only_present)\n    return loss\n\n\ndef lovasz_softmax_flat(probas, labels, only_present=False):\n    \"\"\"\n    Multi-class Lovasz-Softmax loss\n      probas: [P, C] Variable, class probabilities at each prediction (between 0 and 1)\n      labels: [P] Tensor, ground truth labels (between 0 and C - 1)\n      only_present: average only on classes present in ground truth\n    \"\"\"\n    C = probas.size(1)\n    losses = []\n    for c in range(C):\n        fg = (labels == c).float() # foreground for class c\n        if only_present and fg.sum() == 0:\n            continue\n        errors = (Variable(fg) - probas[:, c]).abs()\n        errors_sorted, perm = torch.sort(errors, 0, descending=True)\n        perm = perm.data\n        fg_sorted = fg[perm]\n        losses.append(torch.dot(errors_sorted, Variable(lovasz_grad(fg_sorted))))\n    return f_mean(losses)\n\n\ndef flatten_probas(probas, labels, ignore=None):\n    \"\"\"\n    Flattens predictions in the batch\n    \"\"\"\n    B, C, H, W = probas.size()\n    probas = probas.permute(0, 2, 3, 1).contiguous().view(-1, C)  # B * H * W, C = P, C\n    labels = labels.view(-1)\n    if ignore is None:\n        return probas, labels\n    valid = (labels != ignore)\n    vprobas = probas[valid.nonzero().squeeze()]\n    vlabels = labels[valid]\n    return vprobas, vlabels\n\ndef xloss(logits, labels, ignore=None):\n    \"\"\"\n    Cross entropy loss\n    \"\"\"\n    return F.cross_entropy(logits, Variable(labels), ignore_index=255)\n\n\n# --------------------------- HELPER FUNCTIONS ---------------------------\n\ndef f_mean(l, ignore_nan=False, empty=0):\n    \"\"\"\n    nanmean compatible with generators.\n    \"\"\"\n    l = iter(l)\n    if ignore_nan:\n        l = ifilterfalse(np.isnan, l)\n    try:\n        n = 1\n        acc = next(l)\n    except StopIteration:\n        if empty == 'raise':\n            raise ValueError('Empty mean')\n        return empty\n    for n, v in enumerate(l, 2):\n        acc += v\n    if n == 1:\n        return acc\n    return acc / n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-18T10:14:37.391874Z","iopub.execute_input":"2023-07-18T10:14:37.392699Z","iopub.status.idle":"2023-07-18T10:14:41.270885Z","shell.execute_reply.started":"2023-07-18T10:14:37.392671Z","shell.execute_reply":"2023-07-18T10:14:41.269809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%reload_ext autoreload\n%autoreload 2\n%matplotlib inline\n\nfrom fastai.vision.all import *\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\nimport numpy as np\nimport os\nimport cv2\nimport gc\nimport random\nfrom albumentations import *\nfrom sklearn.model_selection import KFold\nimport matplotlib.pyplot as plt\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2023-07-18T10:18:16.72065Z","iopub.execute_input":"2023-07-18T10:18:16.721254Z","iopub.status.idle":"2023-07-18T10:18:24.812098Z","shell.execute_reply.started":"2023-07-18T10:18:16.721215Z","shell.execute_reply":"2023-07-18T10:18:24.811111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bs = 32\nnfolds = 3\nfold = 0\nSEED = 24\nTRAIN = '/kaggle/input/seg256/train/'\nMASKS = '/kaggle/input/seg256/masks/'\nLABELS = '/kaggle/input/seg256/labels.csv'\nNUM_WORKERS = 4","metadata":{"execution":{"iopub.status.busy":"2023-07-18T12:14:41.310704Z","iopub.execute_input":"2023-07-18T12:14:41.311201Z","iopub.status.idle":"2023-07-18T12:14:41.317519Z","shell.execute_reply.started":"2023-07-18T12:14:41.311159Z","shell.execute_reply":"2023-07-18T12:14:41.316606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    #the following line gives ~10% speedup\n    #but may lead to some stochasticity in the results \n    torch.backends.cudnn.benchmark = True\n    \nseed_everything(SEED)","metadata":{"execution":{"iopub.status.busy":"2023-07-18T10:19:54.998022Z","iopub.execute_input":"2023-07-18T10:19:54.998716Z","iopub.status.idle":"2023-07-18T10:19:55.079948Z","shell.execute_reply.started":"2023-07-18T10:19:54.998681Z","shell.execute_reply":"2023-07-18T10:19:55.078867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://www.kaggle.com/datasets/thedevastator/hubmap-2022-256x256\nmean = np.array([0.7720342, 0.74582646, 0.76392896])\nstd = np.array([0.24745085, 0.26182273, 0.25782376])\n\ndef img2tensor(img,dtype:np.dtype=np.float32):\n    if img.ndim==2 : img = np.expand_dims(img,2)\n    img = np.transpose(img,(2,0,1))\n    return torch.from_numpy(img.astype(dtype, copy=False))\n\nclass HuBMAPDataset(Dataset):\n    def __init__(self, fold=fold, train=True, tfms=None):\n        ids = pd.read_csv(LABELS).id.astype(str).values\n        kf = KFold(n_splits=nfolds,random_state=SEED,shuffle=True)\n        ids = set(ids[list(kf.split(ids))[fold][0 if train else 1]])\n        self.fnames = [fname for fname in os.listdir(TRAIN) if fname.split('_')[0] in ids]\n        self.train = train\n        self.tfms = tfms\n        \n    def __len__(self):\n        return len(self.fnames)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        img = cv2.cvtColor(cv2.imread(os.path.join(TRAIN,fname)), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread(os.path.join(MASKS,fname),cv2.IMREAD_GRAYSCALE)\n        if self.tfms is not None:\n            augmented = self.tfms(image=img,mask=mask)\n            img,mask = augmented['image'],augmented['mask']\n        return img2tensor((img/255.0 - mean)/std),img2tensor(mask)\n    \ndef get_aug(p=1.0):\n    return Compose([\n        HorizontalFlip(),\n        VerticalFlip(),\n        RandomRotate90(),\n        ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, \n                         border_mode=cv2.BORDER_REFLECT),\n        OneOf([\n            OpticalDistortion(p=0.3),\n            GridDistortion(p=.1),\n            PiecewiseAffine(p=0.3),\n        ], p=0.3),\n        OneOf([\n            HueSaturationValue(10,15,10),\n            CLAHE(clip_limit=2),\n            RandomBrightnessContrast(),            \n        ], p=0.3),\n    ], p=p)","metadata":{"execution":{"iopub.status.busy":"2023-07-18T12:16:02.571584Z","iopub.execute_input":"2023-07-18T12:16:02.572017Z","iopub.status.idle":"2023-07-18T12:16:02.591338Z","shell.execute_reply.started":"2023-07-18T12:16:02.571981Z","shell.execute_reply":"2023-07-18T12:16:02.590241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = HuBMAPDataset(tfms=get_aug())\ndl = DataLoader(ds,batch_size=64,shuffle=False,num_workers=NUM_WORKERS)\nimgs,masks = next(iter(dl))\n\nplt.figure(figsize=(16,16))\nfor i,(img,mask) in enumerate(zip(imgs,masks)):\n    img = ((img.permute(1,2,0)*std + mean)*255.0).numpy().astype(np.uint8)\n    plt.subplot(8,8,i+1)\n    plt.imshow(img,vmin=0,vmax=255)\n    plt.imshow(mask.squeeze().numpy(), alpha=0.2)\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    \ndel ds,dl,imgs,masks","metadata":{"execution":{"iopub.status.busy":"2023-07-18T10:45:33.283435Z","iopub.execute_input":"2023-07-18T10:45:33.284232Z","iopub.status.idle":"2023-07-18T10:45:47.126901Z","shell.execute_reply.started":"2023-07-18T10:45:33.284193Z","shell.execute_reply":"2023-07-18T10:45:47.122365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install segmentation_models_pytorch","metadata":{"execution":{"iopub.status.busy":"2023-07-18T12:07:39.338065Z","iopub.execute_input":"2023-07-18T12:07:39.338535Z","iopub.status.idle":"2023-07-18T12:08:04.44362Z","shell.execute_reply.started":"2023-07-18T12:07:39.338493Z","shell.execute_reply":"2023-07-18T12:08:04.442398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport os\nimport random\nimport time\nimport warnings\nwarnings.simplefilter(\"ignore\")\n\n#import pdb\n#import zipfile\n#import pydicom\nfrom albumentations import *\nfrom albumentations.pytorch import ToTensorV2\nimport cv2\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image, ImageFilter\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import KFold\nimport tifffile as tiff\nimport torch\nimport torch.backends.cudnn as cudnn\nimport torch.nn as nn\nfrom torch.nn import functional as F\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset, sampler\nfrom tqdm import tqdm_notebook as tqdm\n\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2023-07-18T12:09:39.702112Z","iopub.execute_input":"2023-07-18T12:09:39.702481Z","iopub.status.idle":"2023-07-18T12:09:42.31172Z","shell.execute_reply.started":"2023-07-18T12:09:39.702448Z","shell.execute_reply":"2023-07-18T12:09:42.310684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_UnetPlusPlus():\n    model =  smp.UnetPlusPlus(\n                 encoder_name='efficientnet-b3',\n                 encoder_weights='imagenet',\n                 in_channels=3,\n                 classes=1)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-07-18T12:10:43.213945Z","iopub.execute_input":"2023-07-18T12:10:43.214483Z","iopub.status.idle":"2023-07-18T12:10:43.220883Z","shell.execute_reply.started":"2023-07-18T12:10:43.214439Z","shell.execute_reply":"2023-07-18T12:10:43.219798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(DiceLoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        \n        #comment out if your model contains a sigmoid or equivalent activation layer\n        inputs = F.sigmoid(inputs)       \n        \n        #flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()                            \n        dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        \n        return 1 - dice","metadata":{"execution":{"iopub.status.busy":"2023-07-18T12:12:24.413568Z","iopub.execute_input":"2023-07-18T12:12:24.414041Z","iopub.status.idle":"2023-07-18T12:12:24.424967Z","shell.execute_reply.started":"2023-07-18T12:12:24.414002Z","shell.execute_reply":"2023-07-18T12:12:24.423949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cv_score = 0\nBATCH_SIZE = 16\nEPOCHS = 10\nNUM_WORKERS = 4\nnfolds = 2\nSEED = 2020\nTH = 0.39 \nDEVICE = ('cuda' if torch.cuda.is_available() else 'cpu')\nfor fold in range(nfolds):\n    ds_t = HuBMAPDataset(fold=fold, train=True, tfms=get_aug())\n    ds_v = HuBMAPDataset(fold=fold, train=False)\n    dataloader_t = torch.utils.data.DataLoader(ds_t,batch_size=BATCH_SIZE, shuffle=False,num_workers=NUM_WORKERS)\n    dataloader_v = torch.utils.data.DataLoader(ds_t,batch_size=BATCH_SIZE, shuffle=False,num_workers=NUM_WORKERS)\n    model = get_UnetPlusPlus().to(DEVICE)\n    \n    optimizer = torch.optim.Adam([\n        {'params': model.decoder.parameters(), 'lr': 1e-3}, \n        {'params': model.encoder.parameters(), 'lr': 1e-3},  \n    ])\n    scheduler = optim.lr_scheduler.OneCycleLR(optimizer=optimizer, pct_start=0.1, div_factor=1e3, \n                                              max_lr=1e-2, epochs=EPOCHS, steps_per_epoch=len(dataloader_t))\n    \n    diceloss = DiceLoss()\n    \n    print(f\"########FOLD: {fold}##############\")\n    \n    for epoch in tqdm(range(EPOCHS)):\n        ###Train\n        model.train()\n        train_loss = 0\n    \n        for data in dataloader_t:\n            optimizer.zero_grad()\n            img, mask = data\n            img = img.to(DEVICE)\n            mask = mask.to(DEVICE)\n        \n            outputs = model(img)\n    \n            loss = diceloss(outputs, mask)\n            loss.backward()\n            optimizer.step()\n            scheduler.step()\n            \n            train_loss += loss.item()\n        train_loss /= len(dataloader_t)\n        \n        print(f\"FOLD: {fold}, EPOCH: {epoch + 1}, train_loss: {train_loss}\")\n        \n        ###Validation\n        model.eval()\n        valid_loss = 0\n        \n        for data in dataloader_v:\n            img, mask = data\n            img = img.to(DEVICE)\n            mask = mask.to(DEVICE)\n        \n            outputs = model(img)\n    \n            loss = diceloss(outputs, mask)\n        \n            valid_loss += loss.item()\n        valid_loss /= len(dataloader_v)\n        \n        print(f\"FOLD: {fold}, EPOCH: {epoch + 1}, valid_loss: {valid_loss}\")\n        \n        \n    ###Save model\n    torch.save(model.state_dict(), f\"FOLD{fold}_.pth\")\n    \n    cv_score += valid_loss\n    \ncv_score = cv_score/nfolds","metadata":{"execution":{"iopub.status.busy":"2023-07-18T12:17:29.287569Z","iopub.execute_input":"2023-07-18T12:17:29.288006Z","iopub.status.idle":"2023-07-18T12:32:11.338406Z","shell.execute_reply.started":"2023-07-18T12:17:29.287966Z","shell.execute_reply":"2023-07-18T12:32:11.335132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"CV score is: {cv_score}\")","metadata":{"execution":{"iopub.status.busy":"2023-07-18T12:34:58.748768Z","iopub.execute_input":"2023-07-18T12:34:58.749326Z","iopub.status.idle":"2023-07-18T12:34:58.755112Z","shell.execute_reply.started":"2023-07-18T12:34:58.749283Z","shell.execute_reply":"2023-07-18T12:34:58.754143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}