{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":7317141,"sourceType":"datasetVersion","datasetId":4246147},{"sourceId":153879438,"sourceType":"kernelVersion"}],"dockerImageVersionId":30627,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#Surface Dice Metric\n#https://www.kaggle.com/code/metric/surface-dice-metric/notebook\n%run /kaggle/input/surface-dice-metric-notebook/surface-dice-metric.ipynb","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:31.048186Z","iopub.execute_input":"2024-01-31T10:55:31.048548Z","iopub.status.idle":"2024-01-31T10:55:31.345966Z","shell.execute_reply.started":"2024-01-31T10:55:31.048520Z","shell.execute_reply":"2024-01-31T10:55:31.345030Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BB_PATH = 'nvidia/segformer-b1-finetuned-cityscapes-1024-1024'","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:31.347771Z","iopub.execute_input":"2024-01-31T10:55:31.348044Z","iopub.status.idle":"2024-01-31T10:55:31.351986Z","shell.execute_reply.started":"2024-01-31T10:55:31.348019Z","shell.execute_reply":"2024-01-31T10:55:31.351032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import glob\n\ntrain_image_paths = list(glob.glob('/kaggle/input/blood-vessel-segmentation/train/*/images/*.tif'))\ntrain_image_paths.sort()\n\ntrain_label_paths = list(glob.glob('/kaggle/input/blood-vessel-segmentation/train/*/labels/*.tif'))\ntrain_label_paths.sort()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-01-31T10:55:31.353067Z","iopub.execute_input":"2024-01-31T10:55:31.353394Z","iopub.status.idle":"2024-01-31T10:55:33.098751Z","shell.execute_reply.started":"2024-01-31T10:55:31.353357Z","shell.execute_reply":"2024-01-31T10:55:33.097742Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#https://huggingface.co/docs/transformers/tasks/semantic_segmentation\nfrom datasets import Dataset, DatasetDict, Image\n\ndef create_dataset(image_paths, label_paths):\n    dataset = Dataset.from_dict({\"image\": sorted(image_paths),\n                                \"label\": sorted(label_paths)})\n    dataset = dataset.cast_column(\"image\", Image(decode=True))\n    dataset = dataset.cast_column(\"label\", Image(decode=True))\n    return dataset\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:33.101290Z","iopub.execute_input":"2024-01-31T10:55:33.102071Z","iopub.status.idle":"2024-01-31T10:55:33.108204Z","shell.execute_reply.started":"2024-01-31T10:55:33.102024Z","shell.execute_reply":"2024-01-31T10:55:33.107222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision.transforms import Resize, Compose, Grayscale\n\ntransform = Compose([Resize([512, 512])])","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:33.109667Z","iopub.execute_input":"2024-01-31T10:55:33.109990Z","iopub.status.idle":"2024-01-31T10:55:33.121065Z","shell.execute_reply.started":"2024-01-31T10:55:33.109959Z","shell.execute_reply":"2024-01-31T10:55:33.120080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef convert_image(image):\n    # Convert to 8-bit grayscale\n    if image.mode == 'I;16':\n        image = image.convert('RGB')\n\n    return image\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:33.122417Z","iopub.execute_input":"2024-01-31T10:55:33.122747Z","iopub.status.idle":"2024-01-31T10:55:33.131228Z","shell.execute_reply.started":"2024-01-31T10:55:33.122716Z","shell.execute_reply":"2024-01-31T10:55:33.130300Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def transforms(example_batch):\n    example_batch[\"image\"] = [torch.tensor(np.array(transform(convert_image(x)))) for x in example_batch[\"image\"]]\n    example_batch[\"label\"] = [torch.tensor(np.array(transform(convert_image(x))), dtype=torch.long) for x in example_batch[\"label\"]]\n    return example_batch","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:33.132682Z","iopub.execute_input":"2024-01-31T10:55:33.133029Z","iopub.status.idle":"2024-01-31T10:55:33.143443Z","shell.execute_reply.started":"2024-01-31T10:55:33.132998Z","shell.execute_reply":"2024-01-31T10:55:33.142515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = create_dataset(train_image_paths, train_image_paths).shuffle(seed=0).train_test_split(0.1)","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:33.144873Z","iopub.execute_input":"2024-01-31T10:55:33.145249Z","iopub.status.idle":"2024-01-31T10:55:33.284709Z","shell.execute_reply.started":"2024-01-31T10:55:33.145216Z","shell.execute_reply":"2024-01-31T10:55:33.283803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds.set_transform(transforms)","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:33.285956Z","iopub.execute_input":"2024-01-31T10:55:33.286331Z","iopub.status.idle":"2024-01-31T10:55:33.294570Z","shell.execute_reply.started":"2024-01-31T10:55:33.286292Z","shell.execute_reply":"2024-01-31T10:55:33.293756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#https://www.kaggle.com/code/jeanlucvanlite/training-6\nimport torch.nn as nn\nclass 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        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()                            \n        dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)          \n        return 1 - dice\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:33.297934Z","iopub.execute_input":"2024-01-31T10:55:33.298269Z","iopub.status.idle":"2024-01-31T10:55:33.306603Z","shell.execute_reply.started":"2024-01-31T10:55:33.298238Z","shell.execute_reply":"2024-01-31T10:55:33.305621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom transformers import AutoModelForSemanticSegmentation\nfrom transformers import TrainingArguments, Trainer, DefaultDataCollator\nimport torch.nn as nn\n\nclass SemanticSegmentationModelWrapper(nn.Module):\n    def __init__(self, model_path):\n        super(SemanticSegmentationModelWrapper, self).__init__()\n        self.model = AutoModelForSemanticSegmentation.from_pretrained(model_path)\n        self.model.decode_head.classifier = nn.Conv2d(256, 1, kernel_size=(1, 1), stride=(1, 1))\n        self.criterion = DiceLoss()\n\n    def forward(self, image, label=None):\n        outputs = self.model(image.permute(0, -1, 1, 2).float())\n        predictions = nn.functional.interpolate(\n            outputs.logits,\n            size=label.permute(0, -1, 1, 2).shape[-2:],\n            mode=\"bilinear\",\n            align_corners=False,\n        ).sigmoid()\n\n        # Calculate loss if label masks are provided\n        loss = None\n        if label is not None:\n            loss = self.criterion(predictions, (label[:, :, :, 0]/255))\n            #loss.requires_grad=outputs.logits.requires_grad\n\n        return loss, predictions","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:33.308111Z","iopub.execute_input":"2024-01-31T10:55:33.308422Z","iopub.status.idle":"2024-01-31T10:55:33.320562Z","shell.execute_reply.started":"2024-01-31T10:55:33.308393Z","shell.execute_reply":"2024-01-31T10:55:33.319614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = SemanticSegmentationModelWrapper(BB_PATH)","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:33.321788Z","iopub.execute_input":"2024-01-31T10:55:33.322089Z","iopub.status.idle":"2024-01-31T10:55:34.570153Z","shell.execute_reply.started":"2024-01-31T10:55:33.322056Z","shell.execute_reply":"2024-01-31T10:55:34.569122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import vmap\nfrom tensorflow import map_fn\nimport tensorflow as tf\n\n\ndef single_instance_compute_metrics(x, spacing_mm=(1.0, 1.0)):\n    pred, label = np.split(x.numpy(), 2, 0)\n    surface_distances = compute_surface_distances((label > 0.5).squeeze(0), (pred > 0.5).squeeze(0), spacing_mm)\n    surface_dice = compute_surface_dice_at_tolerance(surface_distances, 0.0)\n    return surface_dice\n\n\ndef compute_metrics(x):\n    pred, label = x\n    pred = pred\n    \n    label = (label[:, :, :, 0]!=0).astype(np.uint8)\n    print(pred.shape, label.shape)\n    _x = np.stack([pred.squeeze(1), label], 1)\n    spacing_mm = (1.0, 1.0) \n    surface_dice = map_fn(single_instance_compute_metrics, tf.convert_to_tensor(_x))\n    return {\"surface_dice\": surface_dice.numpy().mean()}\n\n","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:34.571597Z","iopub.execute_input":"2024-01-31T10:55:34.571899Z","iopub.status.idle":"2024-01-31T10:55:34.580305Z","shell.execute_reply.started":"2024-01-31T10:55:34.571874Z","shell.execute_reply":"2024-01-31T10:55:34.579034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\nimport numpy as np\nimport torch\nfrom torch.utils.data._utils.collate import default_collate\n\n#ChatGpt\n\nclass CustomDataCollator:\n    def __call__(self, batch):\n        \"\"\"\n        Custom collate function for handling TiffImageFiles.\n        \"\"\"\n        # Convert TiffImageFiles to tensors\n        batch = [{k: (self.convert_to_tensor(transform(convert_image(v))) if self.is_tiff_image_file(v) else v) for k, v in item.items()} for item in batch]\n\n        # Use the default collate function for the rest of the processing\n        return default_collate(batch)\n\n    @staticmethod\n    def is_tiff_image_file(item):\n        \"\"\"\n        Check if the item is a TiffImageFile.\n        \"\"\"\n        return isinstance(item, Image.Image) and item.format == \"TIFF\"\n\n    @staticmethod\n    def convert_to_tensor(image):\n        \"\"\"\n        Convert a TiffImageFile to a PyTorch tensor.\n        \"\"\"\n        # Convert TiffImageFile to a numpy array\n        image_array = np.array(image)\n\n        # Convert numpy array to an appropriate type (e.g., float32)\n        if image_array.dtype == np.uint16:\n            # Normalize and convert to float32, or convert to uint8 without normalization\n            # Normalize to [0, 1] range if it's a typical image data\n            image_array = image_array.astype(np.float32)# / 65535.0\n            # Or convert to uint8 without normalization\n            # image_array = (image_array / 256).astype(np.uint8)\n\n        return torch.tensor(image_array, dtype=torch.float32)","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:55:34.581874Z","iopub.execute_input":"2024-01-31T10:55:34.582238Z","iopub.status.idle":"2024-01-31T10:55:34.595021Z","shell.execute_reply.started":"2024-01-31T10:55:34.582206Z","shell.execute_reply":"2024-01-31T10:55:34.594147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"training_args = TrainingArguments(\n    output_dir=\"segformer-b0-scene-parse-150\",\n    learning_rate=6e-5,\n    num_train_epochs=2,\n    per_device_train_batch_size=16,\n    per_device_eval_batch_size=32,\n    save_total_limit=10,\n    evaluation_strategy=\"steps\",\n    save_strategy=\"steps\",\n    save_steps=200,\n    eval_steps=10,\n    logging_steps=1,\n    eval_accumulation_steps=5,\n    remove_unused_columns=False,\n    report_to='none'\n)\n\ntrainer = Trainer(\n    model = model,\n    args = training_args,\n    train_dataset = ds['train'],\n    eval_dataset = ds['test'],\n    #compute_metrics = compute_metrics,\n    data_collator = CustomDataCollator()\n)\n\ntrainer.train()","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:59:03.444286Z","iopub.execute_input":"2024-01-31T10:59:03.444644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model, '/kaggle/working/segformer-b1-finetuned-cityscapes-1024-1024-2e512.pt')","metadata":{"execution":{"iopub.status.busy":"2024-01-31T10:59:01.318984Z","iopub.status.idle":"2024-01-31T10:59:01.319325Z","shell.execute_reply.started":"2024-01-31T10:59:01.319167Z","shell.execute_reply":"2024-01-31T10:59:01.319182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}