{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q segmentation_models_pytorch\n!pip install -q scikit-learn==1.0\nimport sys\nsys.path.append(\"../input/monai-v060-deep-learning-in-healthcare-imaging\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    IMAGE_SIZE = 256\n    TRAIN_BATCH_SIZE = 2\n    VALID_BATCH_SIZE = 2*TRAIN_BATCH_SIZE\n    EPOCHS = 10\n    DEVICE = \"cuda\"\n    OUTPUT = \"/hubmap/\"\n    ENCODER = \"resnet18\"\n    ENCODER_WEIGHTS = \"imagenet\"\n    FOLDS = 5\n    LR = 1e-6","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport albumentations as A\nimport cv2\nimport numpy as np\nfrom skimage import io","metadata":{"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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from monai.losses import DiceLoss\nfrom monai.metrics import DiceMetric\n\ndice_metric = DiceMetric(include_background=True, reduction=\"mean\", get_not_nans=False)\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 = DiceLoss(sigmoid=True)(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 DICE 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 +=  DiceLoss(sigmoid=True)(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 DICE loss is {val_loss}')\n        \n    return dice_score","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import monai\nmodel = monai.networks.nets.UNet(\n    spatial_dims=2,\n    in_channels=3,\n    out_channels=1,\n    channels=(16, 32, 64, 128, 256),\n    strides=(2, 2, 2, 2),\n    num_res_units=2,\n).to(DEVICE)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\ndf = pd.read_csv(\"../input/hubmap-folds/train_256x256_5folds.csv\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_transforms = {\n    \"train\": A.Compose([\n        A.Resize(config.IMAGE_SIZE,config.IMAGE_SIZE, interpolation=cv2.INTER_NEAREST),\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.VerticalFlip(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=10, p=0.5),\n        A.OneOf([\n            A.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n# #             A.OpticalDistortion(distort_limit=0.05, shift_limit=0.05, p=1.0),\n            A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=1.0)\n        ], p=0.25),\n        A.CoarseDropout(max_holes=8, max_height=config.IMAGE_SIZE//20, max_width=config.IMAGE_SIZE//20,\n                        min_holes=5, fill_value=0, mask_fill_value=0, p=0.5),\n        ], p=1.0),\n    \n    \"valid\": A.Compose([\n        A.Resize(config.IMAGE_SIZE,config.IMAGE_SIZE, interpolation=cv2.INTER_NEAREST),\n        A.Normalize(\n        mean=[0.7720342,  0.74582646, 0.76392896],\n        std=[0.24745085, 0.26182273, 0.25782376],\n        max_pixel_value=255.0,\n        p=1.0,\n    ),    \n    ], p=1.0)\n}","metadata":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scores = []\n\nfor fold in range(5):\n    best_metric = 999999999\n    model.to(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=2,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=4,shuffle=False,pin_memory=True) \n\n    optimizer = torch.optim.Adam(model.parameters(),lr=1e-2)\n    print(f'============================== FOLD -- {fold} ==============================')\n    for epoch in range(2):\n        print(f'==================== Epoch -- {epoch} ====================')\n        train(model=model,train_loader=train_loader,device=DEVICE,optimizer=optimizer)\n\n        dice_score = eval(model=model,valid_loader=valid_loader,device=DEVICE,optimizer=optimizer)\n        \n        print(f'DICE Metric={dice_score}')\n\n#         if best_metric >= dice_metric:\n#             best_metric = dice_metric\n#             torch.save(model.state_dict(),f'model-epoch-{epoch}'+str(fold)+'.pth')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}