{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":9959,"sourceType":"datasetVersion","datasetId":6885},{"sourceId":171346727,"sourceType":"kernelVersion"},{"sourceId":172770054,"sourceType":"kernelVersion"},{"sourceId":176715899,"sourceType":"kernelVersion"},{"sourceId":176716968,"sourceType":"kernelVersion"},{"sourceId":176721026,"sourceType":"kernelVersion"}],"dockerImageVersionId":30673,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd \n\nimport os\n\nfrom encode import rle_encode\nfrom vision_models import ResNetUNet\n\nfrom dataset import ImageDataset, train_test_valid_split, broadcast, next_remainder\nfrom dataset import gray_to_tiled_tensor, image_feature_tensor, label_feature_tensor\nfrom torch.utils.data import Subset, DataLoader\nimport glob\nfrom matplotlib import pyplot as plt \nimport torch\nfrom torch.optim import lr_scheduler, Adam\nimport copy\nimport time\nimport torch.nn.functional as F\nfrom  losses import DiceLoss, IOU\nfrom collections import defaultdict\nimport skimage\nimport gc\nimport subprocess\nfrom metric import score\nfrom functools import partial\nfrom metric import compute_surface_dice_at_tolerance, compute_surface_distances","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2024-05-10T19:09:50.750984Z","iopub.status.idle":"2024-05-10T19:09:50.751344Z","shell.execute_reply.started":"2024-05-10T19:09:50.751159Z","shell.execute_reply":"2024-05-10T19:09:50.751173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"offline_mode = True","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.752586Z","iopub.status.idle":"2024-05-10T19:09:50.752885Z","shell.execute_reply.started":"2024-05-10T19:09:50.752738Z","shell.execute_reply":"2024-05-10T19:09:50.752751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_resnet18_weigths():\n    if not os.path.exists('/root/.cache/torch/hub/checkpoints/'):\n        os.makedirs('/root/.cache/torch/hub/checkpoints/')\n    source_path = \"/kaggle/input/resnet18/resnet18.pth\"\n    destination_path = \"/root/.cache/torch/hub/checkpoints/resnet18-f37072fd.pth\"\n    subprocess.run([\"cp\", source_path, destination_path])","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.754010Z","iopub.status.idle":"2024-05-10T19:09:50.754345Z","shell.execute_reply.started":"2024-05-10T19:09:50.754157Z","shell.execute_reply":"2024-05-10T19:09:50.754170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if offline_mode:\n    load_resnet18_weigths()","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.756093Z","iopub.status.idle":"2024-05-10T19:09:50.756558Z","shell.execute_reply.started":"2024-05-10T19:09:50.756327Z","shell.execute_reply":"2024-05-10T19:09:50.756347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_root = \"/kaggle/input\"\n\ntrain_path = data_root  + \"/blood-vessel-segmentation/train/\"\n\nsource_path = train_path + \"kidney_3_sparse/images/*\" \n\ntarget_path = train_path + \"kidney_3_dense/images/\" \ntest_path = data_root + \"/blood-vessel-segmentation/test\"\n\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.757757Z","iopub.status.idle":"2024-05-10T19:09:50.758063Z","shell.execute_reply.started":"2024-05-10T19:09:50.757914Z","shell.execute_reply":"2024-05-10T19:09:50.757926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ngroup_weights = [\n    {\"group\" :\"kidney_1_voi\", \"pos_weigth\":1, \"remove_background\": False, \"dilate_label\": True, \"image_files\":[], \"label_files\":[]},\n    {\"group\" :\"kidney_1_dense\",\"pos_weigth\":1, \"remove_background\": True, \"dilate_label\": False, \"image_files\":[], \"label_files\":[]},\n    {\"group\" :\"kidney_3_dense\",\"pos_weigth\":1, \"remove_background\": False, \"dilate_label\": False, \"image_files\":[], \"label_files\":[]},\n    {\"group\" :\"kidney_2\", \"pos_weigth\":0.85, \"remove_background\": False, \"dilate_label\": False, \"image_files\":[], \"label_files\":[]},\n    {\"group\" :\"kidney_3_sparse\", \"pos_weigth\":0.65, \"remove_background\": False, \"dilate_label\": False, \"image_files\":[], \"label_files\":[]}\n]\ndef get_files(folder):\n    _files = list(filter(os.path.isfile, glob.glob(folder + \"*.tif\")))\n    _files.sort()\n    return _files\n\nfor item in group_weights:\n    if item[\"group\"] != \"kidney_3_dense\":\n        image_files = get_files(train_path + item[\"group\"] + \"/images/\")\n    else:\n        source_path = train_path + \"kidney_3_sparse/images/\"\n        label_folder = train_path + \"kidney_3_dense/labels/\"\n        label_files = glob.glob(os.path.join(label_folder, \"*.tif\"))  \n        image_slices = [os.path.basename(label_file) for label_file in label_files] \n        image_files = []\n        for image_slice in image_slices:\n            image_files.append(os.path.join(source_path, image_slice))\n        image_files.sort()\n    item[\"image_files\"] = image_files\n    label_files = get_files(train_path + item[\"group\"] + \"/labels/\")\n    item[\"label_files\"] = label_files\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.759601Z","iopub.status.idle":"2024-05-10T19:09:50.759927Z","shell.execute_reply.started":"2024-05-10T19:09:50.759769Z","shell.execute_reply":"2024-05-10T19:09:50.759783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_width = 192\ntile_height = 192\nmini_batch_size = 32\nnum_class = 2\n\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.761170Z","iopub.status.idle":"2024-05-10T19:09:50.761528Z","shell.execute_reply.started":"2024-05-10T19:09:50.761364Z","shell.execute_reply":"2024-05-10T19:09:50.761378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def collate_fn(data,dilate_label):\n    im_list = []\n    l_list = []\n    desc_list = []\n    for batch in data: \n        desc , im , l = batch\n        desc_list.append(desc)\n        im_list.append(image_feature_tensor(im,(tile_width,tile_height))) \n        l_list.append(label_feature_tensor(l,(tile_width,tile_height),dilate_label=dilate_label))\n    return desc_list, torch.stack(im_list), torch.stack(l_list)\n\n\n\ndef train_val_kidney_dataloaders(image_files,label_files, batch_size = 1,test_split = 0.2, val_split= 0.1, dilate_label= True, remove_background= False):\n    \n\n    train_idx, test_idx, val_idx = train_test_valid_split(list(range(len(image_files))), test_size=test_split, valid_size=val_split)\n    split_indices = {\"train\": train_idx, \"test\": test_idx, \"val\": val_idx}\n    dataloaders = {} \n    # Iterate over different splits\n    for split, idx in split_indices.items():\n        split_image_files = list(Subset(image_files, idx))\n        split_label_files = list(Subset(label_files, idx))\n\n        dataset = ImageDataset(split_image_files,split_label_files,remove_background=remove_background)\n\n        dataloaders[split] = DataLoader(dataset, batch_size=batch_size,collate_fn=partial(collate_fn,dilate_label=dilate_label),  shuffle=True, num_workers=4,pin_memory=True)\n        \n    return dataloaders\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.763024Z","iopub.status.idle":"2024-05-10T19:09:50.763388Z","shell.execute_reply.started":"2024-05-10T19:09:50.763194Z","shell.execute_reply":"2024-05-10T19:09:50.763208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef collate(input_t, mini_batch_size=5):\n    batch_size, n_rows, n_columns, canals, tile_w , tile_h  = input_t.shape\n    flattened = input_t.reshape(batch_size*n_rows*n_columns,canals,tile_w,tile_h)\n    batch_tiles = batch_size*n_rows*n_columns\n    for i in range(0, batch_tiles, mini_batch_size):\n        start = i\n        end = min(i + mini_batch_size, batch_tiles)\n        yield flattened[start:end]  \n        ","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.764598Z","iopub.status.idle":"2024-05-10T19:09:50.764901Z","shell.execute_reply.started":"2024-05-10T19:09:50.764751Z","shell.execute_reply":"2024-05-10T19:09:50.764764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef uncollate(batch, original_shape, mini_batch_size= 5):\n    batch_size, n_rows, n_columns, canals, tile_w , tile_h  = original_shape\n    flattened = torch.zeros(batch_size*n_rows*n_columns,canals,tile_w,tile_h)\n    for i, mini_batch in enumerate(batch):\n        for ti, tile in enumerate(mini_batch):\n            flattened[i*mini_batch_size + ti] = tile \n    prediction = flattened.reshape(original_shape)\n    return prediction\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.766415Z","iopub.status.idle":"2024-05-10T19:09:50.766728Z","shell.execute_reply.started":"2024-05-10T19:09:50.766579Z","shell.execute_reply":"2024-05-10T19:09:50.766592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nnum_class =2\nmodel = ResNetUNet(num_class).to(device)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.768151Z","iopub.status.idle":"2024-05-10T19:09:50.768480Z","shell.execute_reply.started":"2024-05-10T19:09:50.768322Z","shell.execute_reply":"2024-05-10T19:09:50.768336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_class =2 \ndef calc_loss(pred, target, metrics, pos_weight = 1):\n    # add weights (0.8, 0.2 positive negative) \n    \n    bce = F.binary_cross_entropy_with_logits(pred, target, pos_weight=torch.tensor([pos_weight]).to(device))\n\n    #mse = fl(pred.type(torch.float),target.type(torch.int64))#torch.nn.CrossEntropyLoss()(target,pred)#F.mse_loss(pred, target)\n    \n    pred = F.sigmoid(pred)\n    dice = DiceLoss()\n    iou = IOU()\n    dice_loss = dice(target,pred)*pred.shape[0]\n    iou = iou(target,pred)*pred.shape[0]\n    loss = dice_loss + bce*pred.shape[0]\n    metrics[\"dice\"] += dice_loss.detach()\n    metrics[\"iou\"] += iou.detach()\n    metrics[\"bce\"] += bce.detach()*pred.shape[0]\n    metrics[\"loss\"] +=  loss.detach()\n    return loss\n\n\ndef print_metrics(metrics, epoch_samples, phase):\n    outputs = []\n    for k in metrics.keys():\n        outputs.append(\"{}: {:4f}\".format(k, metrics[k] / epoch_samples))\n    print(\"{}: {}\".format(phase, \", \".join(outputs)))\n\ndef train_model(model, dataloaders, optimizer,pos_weight = 1, num_epochs=1, mini_batch_size = 64):\n    \n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_loss = 1e10\n    \n    losses = {\"train\": [], \"val\": []}\n    scores = []\n    for epoch in range(num_epochs):\n        print('Epoch {}/{}'.format(epoch, num_epochs - 1))\n        print('-' * 10)\n        since = time.time()\n        for phase in ['train', 'val']:\n            if phase == 'train':\n                for param_group in optimizer.param_groups:\n                    print(\"LR\", param_group['lr'])\n                model.train()  # Set model to training mode\n            else:\n                model.eval()  # Set model to evaluate mode\n\n            #inc = 0\n            metrics = defaultdict(float)\n            epoch_samples = 0\n            dl_count = 0\n            \n            for desc, inputs_tensor, labels_tensor in dataloaders[phase]:\n            #for image, label in dataloaders[phase]:\n\n                optimizer.zero_grad()\n                # forward\n                # track history if only in train\n                with torch.set_grad_enabled(phase == 'train'):\n                    inc = 0\n                    # TODO shape label tensor\n\n                    label_tensor_shape = labels_tensor.shape\n                    \n                    mini_batch_predictions= []\n                    # flattened_prediction \n                    for images, labels in zip(collate(inputs_tensor,mini_batch_size), collate(labels_tensor,mini_batch_size)):\n                        outputs = model(images.to(device))\n                        \n                        loss = calc_loss(outputs, labels.to(device), metrics,pos_weight)\n                        losses[phase].append([metrics[\"dice\"].cpu().detach()/epoch_samples,metrics[\"iou\"].cpu().detach()/epoch_samples,metrics[\"loss\"].cpu().detach()/epoch_samples])\n                        if phase == 'train':\n                            optimizer.zero_grad()\n                            loss.backward()\n                            optimizer.step()\n                        \n                        torch.cuda.memory_cached() \n                        epoch_samples += images.size(0)\n                    \n                    epoch_loss = metrics['loss'] / epoch_samples\n                with torch.no_grad():\n                        torch.cuda.empty_cache()\n                        gc.collect()\n                \n                \n                if phase == 'val' and epoch_loss < best_loss:\n                    print(\"saving best model\")\n                    best_loss = epoch_loss\n                    best_model_wts = copy.deepcopy(model.state_dict())\n\n                dl_count +=1         \n        torch.save(best_model_wts, 'model.pth')\n        time_elapsed = time.time() - since\n        print('{:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))\n\n    return losses, model\n\n\ndef apply_otsu(image):\n            \n    low , high = skimage.filters.threshold_multiotsu(image)\n    thresh_mask = skimage.filters.apply_hysteresis_threshold(image,low, high )\n    return thresh_mask\ndef plot_losses(losses,keys = [\"dice\",\"iou\",\"loss\"]):\n    train_losses = np.array(losses[\"train\"])\n    val_losses = np.array(losses[\"val\"])\n    plt.figure()\n    plt.title(\"train loss\")\n\n    \n    for i in range(3):\n        plt.plot(train_losses[:,i],label=\"train \" + keys[i])\n    plt.legend(loc='upper left')\n    plt.figure()\n    plt.title(\"val loss\")\n    for i in range(3):   \n        plt.plot(val_losses[:,i], label= \"validation \" + keys[i])\n    plt.legend(loc='upper left')\n    plt.show()\n    \ndef run(model, num_epochs , image_files, label_files,group , pos_weight =1 , batch_size = 2, mini_batch_size = 64, path_pretrained= None, lr= 1e-4, dilate_label= True, remove_background= False, plot = False):\n\n    optimizer_ft = Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=lr, weight_decay=1e-6)\n\n    #exp_lr_scheduler = lr_scheduler.StepLR(optimizer_ft, step_size=30, gamma=0.1)\n    dataloaders = train_val_kidney_dataloaders(image_files,label_files,batch_size=2, test_split=0.1, val_split=0.4, dilate_label=dilate_label, remove_background=remove_background)\n    if path_pretrained:\n        model.load_state_dict(torch.load(path_pretrained))\n    \n    losses, model = train_model(model,dataloaders, optimizer_ft, pos_weight = pos_weight, num_epochs=num_epochs, mini_batch_size=mini_batch_size)\n    plot_losses(losses)\n    test_dl = dataloaders[\"test\"]\n    print(\"---model score in ----\")\n    print(group)\n    print(\"score = \")\n    print(model_score(model, test_dl, group,plot))\n    return model\ndef model_score(model, dataloader, group , plot= False):\n    threshold =0.5\n    total_dice = 0 \n    total_elements = 0 \n    i = 0 \n    for  data in dataloader:\n        desc, im_t, l_t = data\n        mini_batch_predictions= []\n        for mini_batch in collate(im_t,mini_batch_size):\n            output = model(mini_batch.to(device))\n            pred =  F.sigmoid(output)\n            pred = pred.data.cpu() \n            mini_batch_predictions.append(pred)\n            torch.cuda.empty_cache()\n            gc.collect()\n            torch.cuda.memory_cached()\n        prediction_t = uncollate(mini_batch_predictions,l_t.shape,mini_batch_size) \n        predictions= []\n        n  = im_t.size(0)\n        total_elements += n\n        \n        for j in range(n):\n            canal1= broadcast(prediction_t[j][:,:,0,:,:])\n            canal2= broadcast(prediction_t[j][:,:,0,:,:])\n            image = broadcast(im_t[j][:,:,0,:,:])\n            threshold_prediction = canal1 > threshold\n            mask_gt = broadcast(l_t[j][:,:,0,:,:]).astype(\"bool\")\n            surface_dist = compute_surface_distances(\n                mask_gt=mask_gt.astype(\"bool\"),\n                mask_pred=threshold_prediction.astype(\"bool\"),\n                spacing_mm=(1,1),\n            )\n            dice = compute_surface_dice_at_tolerance(\n                surface_dist,\n                tolerance_mm=0.0,\n            )\n            total_dice += np.nan_to_num(dice)\n            if plot == True and  i <= 2:\n                fig, axs = plt.subplots( 1,4, figsize=(12,4))\n                \n                axs[0].imshow(image)\n                axs[0].axis('off')\n                axs[1].set_title(\"ground truth\")\n                axs[1].imshow(mask_gt)\n                axs[1].axis('off')\n                axs[2].set_title(\"prediction\")\n                im = axs[2].imshow(canal1)\n                axs[3].imshow(canal2)\n                axs[3].axis('off')\n                plt.tight_layout()\n                cbar = fig.colorbar(im, ax=axs, orientation='vertical')\n                cbar.set_label('Intensity') \n        i+=1  \n    return total_dice/total_elements","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.769854Z","iopub.status.idle":"2024-05-10T19:09:50.770175Z","shell.execute_reply.started":"2024-05-10T19:09:50.770019Z","shell.execute_reply":"2024-05-10T19:09:50.770032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nnum_epochs = 1\n\nfor params in group_weights:\n    group = params[\"group\"]\n    print(\"training group \")\n    print(group)\n    files = [train_path + group + \"/images/\",train_path + group + \"/labels/\" ]\n    model = run(model,num_epochs,params[\"image_files\"],params[\"label_files\"], group,params[\"pos_weigth\"], dilate_label=params[\"dilate_label\"], remove_background=params[\"remove_background\"],batch_size=2,mini_batch_size=mini_batch_size,plot= True)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.771211Z","iopub.status.idle":"2024-05-10T19:09:50.771564Z","shell.execute_reply.started":"2024-05-10T19:09:50.771406Z","shell.execute_reply":"2024-05-10T19:09:50.771420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"submission\")","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.773093Z","iopub.status.idle":"2024-05-10T19:09:50.773446Z","shell.execute_reply.started":"2024-05-10T19:09:50.773257Z","shell.execute_reply":"2024-05-10T19:09:50.773271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict(model, image):\n    mini_batch_size = 24\n    w,h = image.shape\n    image_equalized = skimage.exposure.equalize_adapthist(image)\n    im_t = image_feature_tensor(image_equalized*256,(tile_width,tile_height))\n    reshaped_tensor = im_t.reshape(1, *im_t.shape)\n    mini_batch_predictions= []\n    for mini_batch in collate(reshaped_tensor,mini_batch_size):\n        output = model(mini_batch.to(device))\n        pred =  F.sigmoid(output)\n        pred = pred.data.cpu() \n        mini_batch_predictions.append(pred)\n        torch.cuda.empty_cache()\n        gc.collect()\n        torch.cuda.memory_cached()\n    shape = np.array(reshaped_tensor.shape)\n    shape[3] = num_class\n    prediction_t = uncollate(mini_batch_predictions,tuple(shape),mini_batch_size) \n    \n    canal1= broadcast(prediction_t[0][:,:,0,:,:])\n    #canal2= broadcast(prediction_t[0][:,:,0,:,:])\n    return canal1[:w,:h].astype(np.uint8)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:09:50.775211Z","iopub.status.idle":"2024-05-10T19:09:50.775581Z","shell.execute_reply.started":"2024-05-10T19:09:50.775419Z","shell.execute_reply":"2024-05-10T19:09:50.775433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load best model weights \n# ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ids = []\nrle_masks = []\nthreshold = 0.5\ncontents = os.listdir(test_path)\ndataloaders = {}\n\nfor item in os.listdir(test_path):\n    if os.path.isdir(os.path.join(test_path, item)):\n        content_path = os.path.join(test_path, item)\n        image_files = get_files(content_path + \"/images/\")\n        for i, image_file in enumerate(image_files):\n            image = skimage.io.imread(image_file)\n            prediction = predict(model,image)\n            thresh_prediction = (prediction>threshold).astype(\"bool\")\n            slice_number = (os.path.basename(image_file)).split(\".\")[0]#image_file.split('.')[0].split(\"/\")[-1]\n            ids.append(item + \"_\" + (slice_number))\n            encoded_image = rle_encode(thresh_prediction)\n            if encoded_image == \"\":\n                rle_masks.append(\"1 0\")\n            else:\n                rle_masks.append(encoded_image)\n\n            \nsubmission = pd.DataFrame({'id': ids, 'rle': rle_masks})\nsubmission.to_csv(\"submission.csv\",index=False)   ","metadata":{"execution":{"iopub.status.busy":"2024-05-10T19:13:27.055579Z","iopub.execute_input":"2024-05-10T19:13:27.056490Z","iopub.status.idle":"2024-05-10T19:13:27.549211Z","shell.execute_reply.started":"2024-05-10T19:13:27.056457Z","shell.execute_reply":"2024-05-10T19:13:27.548310Z"},"trusted":true},"execution_count":null,"outputs":[]}]}