{"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":"%%capture\n! pip install segmentation-models-pytorch","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-03T01:10:38.724838Z","iopub.execute_input":"2022-07-03T01:10:38.725258Z","iopub.status.idle":"2022-07-03T01:10:54.285419Z","shell.execute_reply.started":"2022-07-03T01:10:38.725172Z","shell.execute_reply":"2022-07-03T01:10:54.284217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Imports\nimport sys\nimport os\nimport torch\nimport pandas as pd\nimport numpy as np\nfrom skimage import io, transform\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:10:54.287792Z","iopub.execute_input":"2022-07-03T01:10:54.288811Z","iopub.status.idle":"2022-07-03T01:11:03.42194Z","shell.execute_reply.started":"2022-07-03T01:10:54.288767Z","shell.execute_reply":"2022-07-03T01:11:03.420185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    IMAGE_SIZE = 256\n    TRAIN_BATCH_SIZE = 12\n    VALID_BATCH_SIZE = 2*TRAIN_BATCH_SIZE\n    EPOCHS = 10\n    DEVICE = \"cuda\"\n    OUTPUT = \"/hubmap/\"\n    ENCODER = \"ef\"\n    ENCODER_WEIGHTS = \"imagenet\"\n    FOLDS = 5\n    LR = 5e-3\n#     TMAX = int(30000/TRAIN_BATCH_SIZE*EPOCHS)+50","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:11:03.424159Z","iopub.execute_input":"2022-07-03T01:11:03.425278Z","iopub.status.idle":"2022-07-03T01:11:03.434577Z","shell.execute_reply.started":"2022-07-03T01:11:03.425234Z","shell.execute_reply":"2022-07-03T01:11:03.433153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rleToMask(rleString,height,width):\n    rows,cols = height,width\n    rleNumbers = [int(numstring) for numstring in rleString.split(' ')]\n    rlePairs = np.array(rleNumbers).reshape(-1,2)\n    img = np.zeros(rows*cols,dtype=np.uint8)\n    for index,length in rlePairs:\n        index -= 1\n        img[index:index+length] = 255\n    img = img.reshape(cols,rows)\n    img = img.T\n    return img","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:11:03.444672Z","iopub.execute_input":"2022-07-03T01:11:03.445576Z","iopub.status.idle":"2022-07-03T01:11:03.459802Z","shell.execute_reply.started":"2022-07-03T01:11:03.445532Z","shell.execute_reply":"2022-07-03T01:11:03.457583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HubDataset(torch.utils.data.Dataset):\n    def __init__(self, image_paths ,mask_paths=None, transforms=None):\n        self.image_paths = image_paths\n        self.mask_paths = mask_paths\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self, item):\n        image = io.imread(self.image_paths[item])\n        \n        if self.mask_paths is not None:\n            mask = io.imread(self.mask_paths[item])\n            mask = mask.reshape(mask.shape[0],mask.shape[1],1)\n        \n            if self.transforms is not None:\n                augmented = self.transforms(image=image, mask=mask)\n                image = augmented[\"image\"]\n                mask = augmented[\"mask\"]\n                \n            image = np.transpose(image, (2,0,1))\n            mask = np.transpose(mask, (2,0,1))\n            return {\n                \"image\": torch.tensor(image, dtype=float),\n                \"mask\": torch.tensor(mask, dtype=float)\n            }\n        else:\n            if self.transforms is not None:\n                augmented = self.transforms(image=image,)\n                image = augmented[\"image\"]\n        \n            image = np.transpose(image, (2,0,1))\n\n            return {\n                \"image\": torch.tensor(image, dtype=float),}","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:11:03.467194Z","iopub.execute_input":"2022-07-03T01:11:03.468049Z","iopub.status.idle":"2022-07-03T01:11:03.494924Z","shell.execute_reply.started":"2022-07-03T01:11:03.467967Z","shell.execute_reply":"2022-07-03T01:11:03.493326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"../input/hubmap-folds/train_256x256_5folds.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:11:03.497255Z","iopub.execute_input":"2022-07-03T01:11:03.498118Z","iopub.status.idle":"2022-07-03T01:11:03.547651Z","shell.execute_reply.started":"2022-07-03T01:11:03.498076Z","shell.execute_reply":"2022-07-03T01:11:03.546368Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = smp.Unet(encoder_name=\"efficientnet-b7\", \n                 encoder_weights=\"imagenet\", \n                 in_channels=3, \n                 classes=1)","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:11:03.553438Z","iopub.execute_input":"2022-07-03T01:11:03.553867Z","iopub.status.idle":"2022-07-03T01:11:19.215527Z","shell.execute_reply.started":"2022-07-03T01:11:03.553824Z","shell.execute_reply":"2022-07-03T01:11:19.214526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nBCELoss     = smp.losses.SoftBCEWithLogitsLoss()\nLovaszLoss  = smp.losses.LovaszLoss(mode='binary', per_image=False)\nTverskyLoss = smp.losses.TverskyLoss(mode='binary', log_loss=False)\n\ndef criterion(y_pred, y_true):\n    return 0.5*BCELoss(y_pred, y_true) + 0.5*TverskyLoss(y_pred, y_true)\n\ndef train(model,train_loader,device,optimizer):\n    model.train()\n    running_train_loss = 0.0\n    for data in train_loader:\n        inputs = data['image']\n        masks = data['mask']\n\n        inputs = inputs.to(device, dtype=torch.float)\n        masks = masks.to(device, dtype=torch.float)\n\n        optimizer.zero_grad()\n        outputs = model(inputs,)\n        loss = criterion(torch.sigmoid(outputs), masks)\n        loss.backward()\n        optimizer.step()\n        running_train_loss +=loss.item()\n        \n    train_loss_value = running_train_loss/len(train_loader)\n    print(f'train custom Loss is {train_loss_value}')\n    \ndef eval(model,valid_loader,device,optimizer):\n    model.eval()\n    running_dice_score = 0.0\n    running_val_loss = 0.0\n    with torch.no_grad():\n        for data in valid_loader:\n            inputs = data['image']\n            masks = data['mask']\n            \n            inputs = inputs.to(device, dtype=torch.float)\n            masks = masks.to(device, dtype=torch.float)\n\n            output = model(inputs,)\n            running_val_loss +=  criterion(output, masks)\n            \n            output = torch.sigmoid(output)\n            running_dice_score += dice_coef(masks, output)\n \n        val_loss = running_val_loss/len(valid_loader) \n        dice_score = running_dice_score/len(valid_loader)\n        \n        print(f'valid custom Loss is {val_loss}')\n        \n    return dice_score","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:11:19.21684Z","iopub.execute_input":"2022-07-03T01:11:19.217467Z","iopub.status.idle":"2022-07-03T01:11:19.232244Z","shell.execute_reply.started":"2022-07-03T01:11:19.217428Z","shell.execute_reply":"2022-07-03T01:11:19.231268Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\n\ndata_transforms = {\n    \"train\": A.Compose([\n        A.Resize(config.IMAGE_SIZE,config.IMAGE_SIZE),\n            A.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n        max_pixel_value=255.0,\n        p=1.0,\n        ),\n        A.HorizontalFlip(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=10, p=0.5),\n        ], p=1.0),\n    \n    \"valid\": A.Compose([\n        A.Resize(config.IMAGE_SIZE,config.IMAGE_SIZE),\n        A.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n        max_pixel_value=255.0,\n        p=1.0,\n    ),    \n    ], p=1.0)\n}","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:11:19.233799Z","iopub.execute_input":"2022-07-03T01:11:19.234185Z","iopub.status.idle":"2022-07-03T01:11:20.138679Z","shell.execute_reply.started":"2022-07-03T01:11:19.234147Z","shell.execute_reply":"2022-07-03T01:11:20.137661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_coef(y_true, y_pred, thr=0.5, dim=(2,3), epsilon=0.001):\n    y_true = y_true.to(torch.float32)\n    y_pred = (y_pred>thr).to(torch.float32)\n    inter = (y_true*y_pred).sum(dim=dim)\n    den = y_true.sum(dim=dim) + y_pred.sum(dim=dim)\n    dice = ((2*inter+epsilon)/(den+epsilon)).mean(dim=(1,0))\n    return dice","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:11:20.142151Z","iopub.execute_input":"2022-07-03T01:11:20.142553Z","iopub.status.idle":"2022-07-03T01:11:20.148983Z","shell.execute_reply.started":"2022-07-03T01:11:20.142508Z","shell.execute_reply":"2022-07-03T01:11:20.148041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scores = []\n\nfor fold in {0, 1, 2, 3, 4}:\n\n    best_metric = 0.0\n    model.to(config.DEVICE)\n    df_train = df[df.fold != fold].reset_index(drop=True)\n    df_valid = df[df.fold == fold].reset_index(drop=True)\n\n    df_train = df_train.drop(columns = 'fold')\n    df_valid = df_valid.drop(columns = 'fold')\n\n    train_images = df_train.image_path.values.tolist()\n    valid_images = df_valid.image_path.values.tolist()\n\n\n    train_masks = df_train.mask_path.values\n    valid_masks = df_valid.mask_path.values\n\n    train_dataset = HubDataset(image_paths=train_images,mask_paths=train_masks,transforms=data_transforms[\"train\"])\n    train_loader = torch.utils.data.DataLoader(train_dataset,batch_size=config.TRAIN_BATCH_SIZE,shuffle=True,pin_memory=True) \n    valid_dataset = HubDataset(image_paths=valid_images, mask_paths=valid_masks,transforms=data_transforms[\"valid\"])\n    valid_loader = torch.utils.data.DataLoader(valid_dataset,batch_size=config.VALID_BATCH_SIZE,shuffle=False,pin_memory=True) \n\n    optimizer = torch.optim.AdamW(model.parameters(),lr=config.LR,amsgrad=True )\n    scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.5, verbose=True)\n\n    print(f'============================== FOLD -- {fold} ==============================')\n    for epoch in range(config.EPOCHS):\n        print(f'==================== Epoch -- {epoch} ====================')\n        train(model=model,train_loader=train_loader,device=config.DEVICE,optimizer=optimizer)\n        scheduler.step()\n        dice_score = eval(model=model,valid_loader=valid_loader,device=config.DEVICE,optimizer=optimizer)\n\n        print(f'DICE Metric={dice_score}')\n\n        \n        torch.save(model.state_dict(),'model-'+str(fold)+ f\"epoch-{epoch}\" +'.pth')","metadata":{"execution":{"iopub.status.busy":"2022-07-03T01:11:20.150789Z","iopub.execute_input":"2022-07-03T01:11:20.151638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for data in train_loader:\n    print(data[\"mask\"].shape)\n    print(data[\"image\"].shape)\n    break","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}