{"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":"# Current environment requires torchvision==0.15.1 and segmentation_models_pytorch\n#!pip install -U torch==2.0.0\n#!pip install -U torchvision==0.15.1    # Necessary for training\n#!pip install segmentation_models_pytorch   # Necessary for the model\n\nkaggle = True\n\nif kaggle: \n    # Pip install the necessary packages \n    ! pip install --no-index --no-deps /kaggle/input/segmentationpackage/wheelhouse/pretrainedmodels-0.7.4-py3-none-any.whl\n    ! pip install --no-index --no-deps /kaggle/input/segmentationpackage/wheelhouse/efficientnet_pytorch-0.7.1-py3-none-any.whl\n    ! pip install --no-index --no-deps /kaggle/input/segmentationpackage/wheelhouse/timm-0.6.12-py3-none-any.whl\n    ! pip install --no-index --no-deps /kaggle/input/segmentationpackage/wheelhouse/segmentation_models_pytorch-0.3.2-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:43.520891Z","iopub.execute_input":"2023-05-15T09:08:43.521866Z","iopub.status.idle":"2023-05-15T09:08:52.522471Z","shell.execute_reply.started":"2023-05-15T09:08:43.521827Z","shell.execute_reply":"2023-05-15T09:08:52.521133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\n\n\nimport torchvision.transforms as T2     \n#import torchvision.transforms.v2 as T2  # if data augmentation is needed\n\nimport segmentation_models_pytorch as smp\n\nimport glob\n\nimport numpy as np\nimport PIL.Image as Image\nfrom matplotlib import pyplot as plt\nimport pandas as pd\n\n\n# Currently unused\n#import random\n#import torchvision.transforms.v2.functional as TF\n\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.527107Z","iopub.execute_input":"2023-05-15T09:08:52.527458Z","iopub.status.idle":"2023-05-15T09:08:52.542628Z","shell.execute_reply.started":"2023-05-15T09:08:52.527422Z","shell.execute_reply":"2023-05-15T09:08:52.541448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    TILE_DEPTH = 5\n    TILE_HEIGHT = 224\n    TILE_WIDTH = 224\n    TILE_STRIDE = 112\n    TILE_PADDING = 28\n\n    Z_MID = 32\n\n    INKLABEL_REGION_THRESHOLD = 0.02 # 0.10\n\n    DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n    TRAIN = False \n\n    if kaggle: \n        BASE_PATH = \"/kaggle/input/\" # for submission\n        OUTPUT_PATH = \"/kaggle/working/\" # for submission\n    else:\n        BASE_PATH = \"./\" # if local\n        OUTPUT_PATH = \"./\" # if local\n\n    DEBUG = False\n\n    METHOD = \"fusion\"\n","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.544423Z","iopub.execute_input":"2023-05-15T09:08:52.544795Z","iopub.status.idle":"2023-05-15T09:08:52.556147Z","shell.execute_reply.started":"2023-05-15T09:08:52.544755Z","shell.execute_reply":"2023-05-15T09:08:52.555042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Define the segmentation model\nclass CustomModel1(nn.Module):\n    def __init__(self,config):\n        super().__init__()\n        self.encoder = smp.Unet(\n            encoder_name=\"efficientnet-b3\", #\"resnext50_32x4d\", \n            #encoder_depth=5,\n            encoder_weights=\"imagenet\",\n            in_channels=config.TILE_DEPTH,\n            classes=1\n        )\n    \n    def forward(self,X):\n        y = self.encoder(X)\n        return y","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.559618Z","iopub.execute_input":"2023-05-15T09:08:52.560321Z","iopub.status.idle":"2023-05-15T09:08:52.567877Z","shell.execute_reply.started":"2023-05-15T09:08:52.560280Z","shell.execute_reply":"2023-05-15T09:08:52.566780Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Helper functions and classes**","metadata":{}},{"cell_type":"code","source":"def prepare_data(config,folder_list):\n\n    # Create a dictionary of fragments. Key: fragment_id, Value: image_stack (subvolume + inklabels)\n    # fragment_dict is a list of size 3.\n    # fragment_dict[0][0] returns fragment id\n    # fragment_dict[0][1] returns image_stack\n    # fragment_dict[0][2] returns inklabels \n\n    # Initialize\n    fragment_dict = np.empty(len(folder_list),object)\n\n    # Get z_start and z_end for subvolume\n    z_start = config.Z_MID - (config.TILE_DEPTH // 2)\n    z_end = config.Z_MID + (config.TILE_DEPTH // 2) + 1\n\n    for fragment_id, fragment_path in enumerate(folder_list): \n        # Get link to images\n        surface_volume_paths = sorted( glob.glob(fragment_path+\"surface_volume/*.tif\") )\n\n        # Only get a subvolume\n        surface_volume_paths = surface_volume_paths[z_start:z_end]\n\n        # Read a single image to get (height,width)\n        image = np.array(Image.open(surface_volume_paths[0]),  dtype=np.float32)/65535.0\n        height,width = image.shape \n\n        # Initialize image_stack\n        image_stack = np.zeros([config.TILE_DEPTH,height,width],dtype=np.float32)\n        \n        # Go over each subvolume slice and stack them\n        for i, filename in enumerate(surface_volume_paths):\n            image = np.array(Image.open(filename),  dtype=np.float32)/65535.0\n            image_stack[i,:,:] = image \n        \n        # Target is the inklabels. Expand dims to match the subvolume dimensions\n        inklabels = np.expand_dims(np.array(Image.open( fragment_path + \"inklabels.png\" ).convert(\"1\"),dtype=np.float32), axis=0)\n\n        # Append image_stack and inklabels\n        image_stack = np.concatenate((image_stack,inklabels),axis=0)\n\n        # Print the dimensions to check\n        print(image_stack.shape)\n\n        # For each fragment, set the image_stack\n        fragment_dict[fragment_id] = image_stack\n    \n    return fragment_dict","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.569808Z","iopub.execute_input":"2023-05-15T09:08:52.570432Z","iopub.status.idle":"2023-05-15T09:08:52.584370Z","shell.execute_reply.started":"2023-05-15T09:08:52.570385Z","shell.execute_reply":"2023-05-15T09:08:52.583204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create dataset class\nclass SubVolumeDataset(Dataset):\n    def __init__(self,config,fragment_dict,transform=None):\n\n        # fragment_dict holds all fragment ids and image_stack (subvolume and inklabels)\n        self.fragment_dict = fragment_dict\n\n\n        self.TILE_HEIGHT_WITHOUT_PADDING = config.TILE_HEIGHT\n        self.TILE_WIDTH_WITHOUT_PADDING = config.TILE_WIDTH\n\n        self.TILE_DEPTH = config.TILE_DEPTH\n        self.TILE_HEIGHT = config.TILE_HEIGHT + config.TILE_PADDING\n        self.TILE_WIDTH = config.TILE_WIDTH + config.TILE_PADDING\n        self.TILE_STRIDE = config.TILE_STRIDE\n        \n        self.Z_MID = config.Z_MID\n\n        self.transform = transform\n\n        # This stores the (fragment_id, tile_x_coordinate, tile_y_coordinate)s \n        self.valid_fragments_pixels = []\n\n        # For each fragment, get the tile positions to be used in training\n        for fragment_id in range(len(self.fragment_dict)): \n\n            # Get the size of an image\n            inklabels = self.fragment_dict[fragment_id][self.TILE_DEPTH,:,:]\n            height, width = inklabels.shape\n\n            # Go over each tile and decide whether to add it to dataset or not\n            for x in range(0,width-self.TILE_WIDTH,self.TILE_STRIDE):\n                for y in range(0,height-self.TILE_HEIGHT,self.TILE_STRIDE):\n                    # Get the tile\n                    tile_inklabels = inklabels[y:y+self.TILE_HEIGHT,x:x+self.TILE_WIDTH] \n\n                    # If the region has inklabel pixels more than a threshold, add it (using its tile cordinates) to the dataset\n                    if np.sum(tile_inklabels) >= config.INKLABEL_REGION_THRESHOLD * (self.TILE_HEIGHT * self.TILE_WIDTH):\n                        self.valid_fragments_pixels.append([fragment_id,x,y])\n                        \n    def __len__(self):\n        return len(self.valid_fragments_pixels)\n    \n\n    def segmentation_transform(self,image_stack):\n        transform = T2.Compose([\n            #T2.RandomRotation(10),\n            T2.RandomCrop((self.TILE_HEIGHT_WITHOUT_PADDING,self.TILE_WIDTH_WITHOUT_PADDING))\n            #T2.CenterCrop((self.TILE_HEIGHT_WITHOUT_PADDING,self.TILE_WIDTH_WITHOUT_PADDING))\n        ])\n        transformed_image_stack = transform(image_stack)\n        return transformed_image_stack\n\n\n    def __getitem__(self,index):\n        \n        # Get the fragment_id from the valid_fragment_pixels\n        fragment_id = self.valid_fragments_pixels[index][0]\n\n        # Get the tile coordinates\n        tile_x = self.valid_fragments_pixels[index][1]\n        tile_y = self.valid_fragments_pixels[index][2]\n\n        # Get the tile_image stack\n        tile_image_stack = self.fragment_dict[fragment_id][:,tile_y:tile_y+self.TILE_HEIGHT,tile_x:tile_x+self.TILE_WIDTH] \n\n        # Data augmentation and transform to tensor here...\n        if self.transform is None:\n            #self.transform = transforms.Compose([transforms.ToTensor()])\n\n            # Convert to tensor\n            tile_image_stack = torch.from_numpy(tile_image_stack)\n            \n            # Apply the transform\n            tile_image_stack = self.segmentation_transform(tile_image_stack)\n\n        else:\n            tile_image_stack = self.transform(tile_image_stack)\n        \n        # Get the tile_subvolume and tile_inklabels\n        tile_subvolume = tile_image_stack[0:self.TILE_DEPTH,:,:]\n        tile_inklabels = tile_image_stack[self.TILE_DEPTH,:,:].unsqueeze(dim=0) # A channel is taken, that dim is dropped. Add it back.\n    \n        return tile_subvolume, tile_inklabels","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.587128Z","iopub.execute_input":"2023-05-15T09:08:52.587524Z","iopub.status.idle":"2023-05-15T09:08:52.607881Z","shell.execute_reply.started":"2023-05-15T09:08:52.587486Z","shell.execute_reply":"2023-05-15T09:08:52.606771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Write a custom loss function\nclass DiceBCEwithLogitsLoss(nn.Module):\n    def __init__(self,alpha=0.5,epsilon=0.1):\n        super().__init__()\n        self.alpha = alpha\n        self.epsilon = epsilon\n    \n    def forward(self,y_pred,y):\n\n        # Apply sigmoid to get values in range [0,1]\n        y_pred = nn.functional.sigmoid(y_pred)\n\n        # Flatten\n        y_pred = y_pred.view(-1)\n        y = y.view(-1)\n\n        intersection = (y_pred * y).sum()\n        dice_score = (2*intersection) / (y_pred.sum() + y.sum() + self.epsilon)\n        dice_loss = 1 - dice_score\n\n        BCE_loss = nn.functional.binary_cross_entropy(y_pred,y,reduction='mean')\n\n        DiceBCE = self.alpha * dice_loss + (1 - self.alpha) * BCE_loss\n\n        return DiceBCE ","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.609624Z","iopub.execute_input":"2023-05-15T09:08:52.610006Z","iopub.status.idle":"2023-05-15T09:08:52.622142Z","shell.execute_reply.started":"2023-05-15T09:08:52.609967Z","shell.execute_reply":"2023-05-15T09:08:52.621149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define a function to train the model for one epoch\ndef train_model(model, optimizer, loss_function, train_loader):\n  \n    # Send model to device\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n\n    # Put model in train mode\n    model.train()\n\n    # For loss calculation\n    total_loss = 0\n    #total_correct = 0\n    dataset_size = len(train_loader.dataset)\n\n    # Go over each batch\n    for X, y in train_loader:\n        X = X.to(device)\n        y = y.to(device)\n\n        y_pred = model(X)\n\n        loss = loss_function(y_pred,y)  \n\n        # Zero the gradients, backpropagate the gradients, update the parameters\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        # Calculate loss\n        total_loss += loss_function(y_pred, y).item()\n        #total_correct += (y_pred.argmax(1) == y).type(torch.float).sum().item()\n\n    print(f\"Training Loss:{total_loss/dataset_size:0.6f}\")\n    #print(f\"Training Loss:{total_loss/dataset_size:0.4f},Training Accuracy:{total_correct/dataset_size:0.4f}\")","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.623699Z","iopub.execute_input":"2023-05-15T09:08:52.624297Z","iopub.status.idle":"2023-05-15T09:08:52.636194Z","shell.execute_reply.started":"2023-05-15T09:08:52.624257Z","shell.execute_reply":"2023-05-15T09:08:52.635197Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_test_data(config,test_folder):\n    \n    # Get z_start and z_end\n    z_start = config.Z_MID - (config.TILE_DEPTH // 2)\n    z_end = config.Z_MID + (config.TILE_DEPTH // 2) + 1\n\n    # Get link to images\n    fragment_path = test_folder\n    surface_volume_paths = sorted( glob.glob(fragment_path+\"surface_volume/*.tif\") )\n    surface_volume_paths = surface_volume_paths[z_start:z_end]\n\n    # Read a single image\n    image = np.array(Image.open(surface_volume_paths[0]),  dtype=np.float32)/65535.0\n    height,width = image.shape \n    image_stack = np.zeros([config.TILE_DEPTH,height,width],dtype=np.float32)\n\n    for i, filename in enumerate(surface_volume_paths):\n        image = np.array(Image.open(filename),  dtype=np.float32)/65535.0\n        image_stack[i,:,:] = image     \n\n    return image_stack","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.637844Z","iopub.execute_input":"2023-05-15T09:08:52.638276Z","iopub.status.idle":"2023-05-15T09:08:52.649762Z","shell.execute_reply.started":"2023-05-15T09:08:52.638234Z","shell.execute_reply":"2023-05-15T09:08:52.648477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_model(config,model,image_stack):\n\n    # Send model to device\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n\n    model.eval()\n\n    # Get the dimensions of image_stack\n    depth, height, width = image_stack.shape\n\n    # Create an empty estimate\n    pred_image = np.zeros((height,width),dtype=np.float32)\n\n    with torch.no_grad():\n        # Go over each tile in the image_stack\n        for x in range(0,width-config.TILE_WIDTH,config.TILE_WIDTH):\n            for y in range(0,height-config.TILE_HEIGHT,config.TILE_HEIGHT):\n\n                # Get a tile\n                X_tile = image_stack[:,y:y+config.TILE_HEIGHT,x:x+config.TILE_WIDTH] \n                X_tile = torch.from_numpy(X_tile).unsqueeze(0)\n                X_tile = X_tile.to(device)\n\n                # Make the prediction\n                y_tile = model(X_tile)\n                y_tile = y_tile.view(-1,config.TILE_HEIGHT,config.TILE_WIDTH)\n                #print(y_tile.shape)\n\n                # Place it in the large result\n                pred_image[y:y+config.TILE_HEIGHT,x:x+config.TILE_WIDTH] = y_tile.squeeze().cpu().numpy()\n        \n\n    return pred_image","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.655341Z","iopub.execute_input":"2023-05-15T09:08:52.655676Z","iopub.status.idle":"2023-05-15T09:08:52.666977Z","shell.execute_reply.started":"2023-05-15T09:08:52.655648Z","shell.execute_reply":"2023-05-15T09:08:52.665778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_fusion_model(config,model,image_stack):\n\n    # Send model to device\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n\n    model.eval()\n\n    # Get the dimensions of image_stack\n    depth, height, width = image_stack.shape\n\n    # Start positions #(start_x, start_y)\n    start_positions = [(0,0),\n                       ( int(config.TILE_WIDTH/2),0 ), \n                       ( 0, int(config.TILE_HEIGHT/2) ),\n                       ( int(config.TILE_WIDTH/2), int(config.TILE_HEIGHT/2) )]\n\n    # Create an empty estimate\n    pred_image = np.zeros((len(start_positions),height,width),dtype=np.float32)\n\n    with torch.no_grad():\n        # for each start position\n        for i, (start_x,start_y) in enumerate(start_positions): \n\n            # Go over each tile in the image_stack\n            for x in range(start_x,width-config.TILE_WIDTH,config.TILE_WIDTH):\n                for y in range(start_y,height-config.TILE_HEIGHT,config.TILE_HEIGHT):\n\n                    # Get a tile\n                    X_tile = image_stack[:,y:y+config.TILE_HEIGHT,x:x+config.TILE_WIDTH] \n                    X_tile = torch.from_numpy(X_tile).unsqueeze(0)\n                    X_tile = X_tile.to(device)\n\n                    # Make the prediction\n                    y_tile = model(X_tile)\n                    y_tile = y_tile.view(-1,config.TILE_HEIGHT,config.TILE_WIDTH)\n                    #print(y_tile.shape)\n\n                    # Place it in the large result\n                    pred_image[i,y:y+config.TILE_HEIGHT,x:x+config.TILE_WIDTH] = y_tile.squeeze().cpu().numpy()\n        \n        # Average\n        pred_image = np.mean(pred_image,axis=0)\n\n    return pred_image","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.668942Z","iopub.execute_input":"2023-05-15T09:08:52.669400Z","iopub.status.idle":"2023-05-15T09:08:52.684184Z","shell.execute_reply.started":"2023-05-15T09:08:52.669362Z","shell.execute_reply":"2023-05-15T09:08:52.683145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Training**","metadata":{}},{"cell_type":"code","source":"if config.TRAIN:\n    # Get list of directories \n    train_folder_list = [config.BASE_PATH + \"vesuvius-challenge/train/1/\", \n                        config.BASE_PATH + \"vesuvius-challenge/train/2/\", \n                        config.BASE_PATH + \"vesuvius-challenge/train/3/\"]\n\n    # Prepare training data\n    train_data_dict = prepare_data(config,train_folder_list)\n","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.687629Z","iopub.execute_input":"2023-05-15T09:08:52.688626Z","iopub.status.idle":"2023-05-15T09:08:52.694518Z","shell.execute_reply.started":"2023-05-15T09:08:52.688587Z","shell.execute_reply":"2023-05-15T09:08:52.693415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check out the images\n\nif config.DEBUG:\n\n    # A subvolume\n    fragment_id = 0\n    print(train_data_dict[fragment_id].shape)\n\n    # There are 3 fragments for training\n    print(len(train_data_dict))\n\n\n    # Check out some subvolume channels\n    fig, ax = plt.subplots(1,5)\n    for i in range(5):\n        ax[i].imshow(train_data_dict[fragment_id][i])\n\n    # Check out a subvolume channel\n    X = train_data_dict[fragment_id][2]\n    plt.figure()\n    plt.imshow(X)\n\n    # Check out inklabels\n    y = train_data_dict[fragment_id][config.TILE_DEPTH]\n    plt.figure()\n    plt.imshow(y)\n","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.696083Z","iopub.execute_input":"2023-05-15T09:08:52.696549Z","iopub.status.idle":"2023-05-15T09:08:52.707059Z","shell.execute_reply.started":"2023-05-15T09:08:52.696509Z","shell.execute_reply":"2023-05-15T09:08:52.705988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.DEBUG:\n\n    # Take a look at a sample\n    sample_id = 0\n    slice_id = 0\n\n    tile_image_stack, tile_inklabels = train_dataset[sample_id]\n\n    image = tile_image_stack[slice_id,:,:].view(config.TILE_HEIGHT,config.TILE_WIDTH).numpy()\n    target = tile_inklabels[0,:,:].view(config.TILE_HEIGHT,config.TILE_WIDTH).numpy()\n\n    plt.imshow(image)\n    plt.figure()\n    plt.imshow(target)\n\n    print(tile_image_stack.shape)\n\n    # Check out the DataLoader\n    BATCH_SIZE = 1\n    train_dataloader = DataLoader(train_dataset,batch_size=BATCH_SIZE)\n    X, y = next(iter(train_dataloader))\n    image = X[0,slice_id,:,:].view(config.TILE_HEIGHT,config.TILE_WIDTH).numpy()\n    target = y[0,0,:,:].view(config.TILE_HEIGHT,config.TILE_WIDTH).numpy()\n\n    plt.figure()\n    plt.imshow(image)\n    plt.figure()\n    plt.imshow(target)\n\n    print(X.shape)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.708674Z","iopub.execute_input":"2023-05-15T09:08:52.709161Z","iopub.status.idle":"2023-05-15T09:08:52.720901Z","shell.execute_reply.started":"2023-05-15T09:08:52.709124Z","shell.execute_reply":"2023-05-15T09:08:52.719789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.TRAIN:\n    # To load a state dictionary\n    # Instantiate a model\n    model1 = CustomModel1(config)   \n    \n    # If load from pretrained model\n    #model1.load_state_dict(torch.load(\"model1_augmented_300epochs.pth\"))","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.722850Z","iopub.execute_input":"2023-05-15T09:08:52.723340Z","iopub.status.idle":"2023-05-15T09:08:52.728838Z","shell.execute_reply.started":"2023-05-15T09:08:52.723302Z","shell.execute_reply":"2023-05-15T09:08:52.727699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.DEBUG:\n\n    # Check out the model summary\n    #print(model1)\n\n    from torchinfo import summary\n    summary(model1,(1,config.TILE_DEPTH,config.TILE_HEIGHT,config.TILE_WIDTH))","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.730569Z","iopub.execute_input":"2023-05-15T09:08:52.731084Z","iopub.status.idle":"2023-05-15T09:08:52.738002Z","shell.execute_reply.started":"2023-05-15T09:08:52.731047Z","shell.execute_reply":"2023-05-15T09:08:52.736735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.DEBUG:\n\n    # Model check with dataset and dataloader\n    BATCH_SIZE = 1\n\n    # Instantiate train_dataset and dataloader\n    train_dataset = SubVolumeDataset(config,train_data_dict)\n    train_dataloader = DataLoader(train_dataset,batch_size=BATCH_SIZE)\n\n    # To load a state dictionary\n    # Instantiate a model\n    model1 = CustomModel1(config)   \n    model1.load_state_dict(torch.load(\"model1_augmented_300epochs.pth\"))\n\n    X, y = next(iter(train_dataloader))\n\n    X = X.to(config.DEVICE)\n    y = y.to(config.DEVICE)\n    model1.to(config.DEVICE)\n\n    y_pred = model1(X)\n\n    slice_id = 0\n    image = X[0,slice_id,:,:].view(config.TILE_HEIGHT,config.TILE_WIDTH).cpu().numpy()\n    target = y[0,0,:,:].view(config.TILE_HEIGHT,config.TILE_WIDTH).cpu().numpy()\n    pred = y_pred[0,0,:,:].view(config.TILE_HEIGHT,config.TILE_WIDTH).detach().cpu().numpy()  # detach because tensor has autograd connection\n\n    plt.figure()\n    plt.imshow(image)\n    plt.figure()\n    plt.imshow(target)\n    plt.figure()\n    plt.imshow(pred)\n","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.740084Z","iopub.execute_input":"2023-05-15T09:08:52.740486Z","iopub.status.idle":"2023-05-15T09:08:52.752470Z","shell.execute_reply.started":"2023-05-15T09:08:52.740449Z","shell.execute_reply":"2023-05-15T09:08:52.751601Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.DEBUG:\n\n    # Check out the loss function\n    #loss_function = torch.nn.CrossEntropyLoss()    # this would be if output has multiple channels, target label coded\n    loss_function = torch.nn.BCEWithLogitsLoss()\n    loss = loss_function(y_pred,y) \n    print(loss)\n\n\n    loss_function = DiceBCEwithLogitsLoss()\n    loss = loss_function(y_pred,y) \n    print(loss)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.753862Z","iopub.execute_input":"2023-05-15T09:08:52.754512Z","iopub.status.idle":"2023-05-15T09:08:52.761812Z","shell.execute_reply.started":"2023-05-15T09:08:52.754474Z","shell.execute_reply":"2023-05-15T09:08:52.761083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train or test the model\n\n\nif config.TRAIN == True:\n\n  # Train the model\n\n  # Loss function\n  loss_function = DiceBCEwithLogitsLoss()\n\n  # Optimizer\n  #optimizer = torch.optim.SGD(model1.parameters(),lr=0.001)\n  optimizer = torch.optim.AdamW(model1.parameters(),lr=0.001)\n\n  BATCH_SIZE = 32\n  \n  train_dataset = SubVolumeDataset(config,train_data_dict)\n  train_dataloader = DataLoader(train_dataset,batch_size=BATCH_SIZE)\n\n  num_epochs = 300\n\n  for epoch in range(num_epochs):\n    print('Epoch:',epoch)\n    # Train loss\n    train_model(model1, optimizer, loss_function, train_dataloader)\n\n  # Save the model\n  #torch.save(model1.state_dict(), \"/kaggle/working/model1_augmented_300epochs.pth\")\n  torch.save(model1,config.BASE_PATH + \"pretrained-model/model_efficientnet_b3_400epochs.pth\")\n\nelse: \n  # Load the model from state dictionary\n  #model1 = CustomModel1(config)\n  #model1.load_state_dict(torch.load(config.BASE_PATH + \"pretrained-model/model1_augmented_300epochs.pth\"))\n\n  # Load the entire model\n  #model1 = torch.load(config.BASE_PATH + \"pretrained-model/model_efficientnet_b3_400epochs.pth\")\n  model1 = torch.load(config.BASE_PATH + \"pretrained-model/model_resnext50_32x4d.pth\")\n","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.763239Z","iopub.execute_input":"2023-05-15T09:08:52.763877Z","iopub.status.idle":"2023-05-15T09:08:52.920471Z","shell.execute_reply.started":"2023-05-15T09:08:52.763840Z","shell.execute_reply":"2023-05-15T09:08:52.919321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.DEBUG:\n\n    # Check out the performance on input data \n\n    # Prepare data to test the model \n    test_folder = config.BASE_PATH + \"vesuvius-challenge/train/1/\"\n    test_image_stack = prepare_test_data(config,test_folder)\n\n\n    # Test the model\n    test_result = test_model(config,model1,test_image_stack)\n\n\n    # Get the mask and display the result \n\n    mask_file = test_folder+\"mask.png\"\n    mask = np.array(Image.open(mask_file),  dtype=np.float32)\n    #plt.imshow(mask)\n\n\n    ground_truth = np.array(Image.open( test_folder + \"inklabels.png\" ).convert(\"1\"), dtype=np.float32)\n\n    test_result = mask * test_result\n\n\n    plt.imshow(test_result)\n\n    threshold = 0.0\n    test_result2 = (test_result > threshold).astype(int) \n\n    plt.figure()\n    plt.imshow(test_result2)\n    plt.title(\"Prediction\")\n\n    plt.figure()\n    plt.imshow(ground_truth)\n    plt.title(\"Ground truth\")","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.922156Z","iopub.execute_input":"2023-05-15T09:08:52.923118Z","iopub.status.idle":"2023-05-15T09:08:52.932890Z","shell.execute_reply.started":"2023-05-15T09:08:52.923075Z","shell.execute_reply":"2023-05-15T09:08:52.931566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Submission**","metadata":{}},{"cell_type":"code","source":"# Submission\ndef rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    # pixels = (pixels >= thr).astype(int)\n    \n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.934561Z","iopub.execute_input":"2023-05-15T09:08:52.935174Z","iopub.status.idle":"2023-05-15T09:08:52.946032Z","shell.execute_reply.started":"2023-05-15T09:08:52.935135Z","shell.execute_reply":"2023-05-15T09:08:52.944949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.METHOD == \"single\":\n\n    base_path = config.BASE_PATH + \"vesuvius-challenge/test/\"\n    fragment_ids = [\"a\",\"b\"]\n\n    results = []\n    for fragment_id in fragment_ids:\n\n        test_folder = base_path + fragment_id + \"/\"\n\n        mask_file = test_folder + \"mask.png\"\n        mask = np.array(Image.open(mask_file),  dtype=np.float32)\n\n\n        # Do the prediction\n        test_image_stack = prepare_test_data(config,test_folder)\n        test_result = test_model(config,model1,test_image_stack)\n        print('result data type',type(test_result[0][0]))\n\n        test_result = mask * test_result\n\n        threshold = 0.0\n        test_result = (test_result > threshold).astype(int) \n\n        plt.figure()\n        plt.imshow(test_result)\n\n\n        # Get the rle\n        inklabels_rle = rle(test_result)\n\n        results.append((fragment_id, inklabels_rle))\n\n\n    sub = pd.DataFrame(results, columns=['Id', 'Predicted'])\n    sub.to_csv(config.OUTPUT_PATH + \"submission.csv\", index=False)\n","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.948883Z","iopub.execute_input":"2023-05-15T09:08:52.949939Z","iopub.status.idle":"2023-05-15T09:08:52.960632Z","shell.execute_reply.started":"2023-05-15T09:08:52.949749Z","shell.execute_reply":"2023-05-15T09:08:52.959283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if config.METHOD == \"fusion\":\n    \n    base_path = config.BASE_PATH + \"vesuvius-challenge/test/\"\n    fragment_ids = [\"a\",\"b\"]\n\n    results = []\n    for fragment_id in fragment_ids:\n\n        test_folder = base_path + fragment_id + \"/\"\n\n        mask_file = test_folder + \"mask.png\"\n        mask = np.array(Image.open(mask_file),  dtype=np.float32)\n\n\n        # Do the prediction\n        test_image_stack = prepare_test_data(config,test_folder)\n        test_result = test_fusion_model(config,model1,test_image_stack)\n        print('result data type',type(test_result[0][0]))\n\n        test_result = mask * test_result\n\n        threshold = 0.0\n        test_result = (test_result > threshold).astype(int) \n\n        plt.figure()\n        plt.imshow(test_result)\n\n\n        # Get the rle\n        inklabels_rle = rle(test_result)\n\n        results.append((fragment_id, inklabels_rle))\n\n\n    sub = pd.DataFrame(results, columns=['Id', 'Predicted'])\n    sub.to_csv(config.OUTPUT_PATH + \"submission.csv\", index=False)\n    ","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:08:52.962606Z","iopub.execute_input":"2023-05-15T09:08:52.963024Z","iopub.status.idle":"2023-05-15T09:09:44.083972Z","shell.execute_reply.started":"2023-05-15T09:08:52.962984Z","shell.execute_reply":"2023-05-15T09:09:44.082784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Ensemble \n\nif config.METHOD == \"ensemble\":\n\n    model_paths = [\"pretrained-model/model_resnext50_32x4d_UNet_600epochs.pth\", \n                   \"pretrained-model/model_resnext50_32x4d_DeepLab_600epochs.pth\",\n                   \"pretrained-model/model_efficientnet_b3_400epochs.pth\",\n                   \"pretrained-model/model_resnext50_32x4d.pth\"]\n\n    base_path = config.BASE_PATH + \"vesuvius-challenge/test/\"\n    fragment_ids = [\"a\",\"b\"]\n\n\n    results = []\n    for fragment_id in fragment_ids:\n\n        test_folder = base_path + fragment_id + \"/\"\n\n        mask_file = test_folder + \"mask.png\"\n        mask = np.array(Image.open(mask_file),  dtype=np.float32)\n\n\n        # Do the prediction\n        test_image_stack = prepare_test_data(config,test_folder)\n\n        test_results = np.zeros((len(model_paths),np.size(mask,0),np.size(mask,1)))\n        for i, model_path in enumerate(model_paths):\n            model1 = torch.load(config.BASE_PATH + model_path)    \n            #test_result = test_fusion_model(config,model1,test_image_stack)\n            test_result = test_model(config,model1,test_image_stack)\n            \n            test_results[i,:,:] = test_result\n\n        print(test_results.shape)\n\n        test_result = np.mean(test_results,axis=0)\n        test_result = mask * test_result\n\n        threshold = 0.0\n        test_result = (test_result > threshold).astype(int) \n\n        plt.figure()\n        plt.imshow(test_result)\n\n\n        # Get the rle\n        inklabels_rle = rle(test_result)\n\n        results.append((fragment_id, inklabels_rle))\n\n    sub = pd.DataFrame(results, columns=['Id', 'Predicted'])\n    sub.to_csv(config.OUTPUT_PATH + \"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T09:09:44.085654Z","iopub.execute_input":"2023-05-15T09:09:44.086327Z","iopub.status.idle":"2023-05-15T09:09:44.099140Z","shell.execute_reply.started":"2023-05-15T09:09:44.086287Z","shell.execute_reply":"2023-05-15T09:09:44.098034Z"},"trusted":true},"execution_count":null,"outputs":[]}]}