{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52279,"databundleVersionId":5822112,"sourceType":"competition"}],"dockerImageVersionId":30512,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"raw","source":"import pytorch_lightning as pl","metadata":{"execution":{"iopub.status.busy":"2023-06-21T06:36:16.787645Z","iopub.execute_input":"2023-06-21T06:36:16.788053Z","iopub.status.idle":"2023-06-21T06:36:27.468509Z","shell.execute_reply.started":"2023-06-21T06:36:16.788002Z","shell.execute_reply":"2023-06-21T06:36:27.467541Z"}}},{"cell_type":"markdown","source":"# Day 1 : Loading the data","metadata":{}},{"cell_type":"code","source":"import os\nimport json\nfrom PIL import Image\nfrom collections import Counter\n\nimport numpy as np\nimport pandas as pd\nimport plotly.express as px\nimport plotly.graph_objects as go\nimport tifffile as tiff\n\n\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\n\nimport cv2\n\n","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:51.657095Z","iopub.execute_input":"2024-08-25T11:33:51.657387Z","iopub.status.idle":"2024-08-25T11:33:52.680985Z","shell.execute_reply.started":"2024-08-25T11:33:51.657362Z","shell.execute_reply":"2024-08-25T11:33:52.680037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open('/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl', 'r') as json_file:\n    json_list = list(json_file)","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:52.683064Z","iopub.execute_input":"2024-08-25T11:33:52.683624Z","iopub.status.idle":"2024-08-25T11:33:53.890055Z","shell.execute_reply.started":"2024-08-25T11:33:52.68359Z","shell.execute_reply":"2024-08-25T11:33:53.88883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tiles_dicts = []\nfor json_str in json_list:\n    tiles_dicts.append(json.loads(json_str))","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:53.891689Z","iopub.execute_input":"2024-08-25T11:33:53.892025Z","iopub.status.idle":"2024-08-25T11:33:58.090889Z","shell.execute_reply.started":"2024-08-25T11:33:53.891995Z","shell.execute_reply":"2024-08-25T11:33:58.089901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# with open('/kaggle/input/hubmap-hacking-the-human-vasculature/polygons.jsonl', 'r') as json_file:\n#     json_list = list(json_file)\n    \n# tiles_dicts = []\n# for json_str in json_list:\n#     tiles_dicts.append(json.loads(json_str))","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:58.091982Z","iopub.execute_input":"2024-08-25T11:33:58.092285Z","iopub.status.idle":"2024-08-25T11:33:58.096492Z","shell.execute_reply.started":"2024-08-25T11:33:58.092261Z","shell.execute_reply":"2024-08-25T11:33:58.09556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile_meta_df = pd.read_csv(\"/kaggle/input/hubmap-hacking-the-human-vasculature/tile_meta.csv\")\ntile_meta_df","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:58.099376Z","iopub.execute_input":"2024-08-25T11:33:58.099691Z","iopub.status.idle":"2024-08-25T11:33:58.191753Z","shell.execute_reply.started":"2024-08-25T11:33:58.099661Z","shell.execute_reply":"2024-08-25T11:33:58.190895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_cartesian_coords(coords, img_height):\n    coords_array = np.array(coords).squeeze()\n    xs = coords_array[:, 0]\n    ys = -coords_array[:, 1] + img_height\n    return xs, ys\n\ndef plot_annotated_image(image_dict, scale_factor: int = 1.0) -> None:\n    #array = tiff.imread(CFG.img_path_template.format(image_dict[\"id\"]))\n    array = tiff.imread(f'/kaggle/input/hubmap-hacking-the-human-vasculature/train/{image_dict[\"id\"]}.tif')\n    \n    img_example = Image.fromarray(array)\n    annotations = image_dict[\"annotations\"]\n    \n    # create figure\n    fig = go.Figure()\n\n    # constants\n    img_width = img_example.size[0]\n    img_height = img_example.size[1]\n    \n\n    # add invisible scatter trace\n    fig.add_trace(\n        go.Scatter(\n            x=[0, img_width],\n            y=[0, img_height],\n            mode=\"markers\",\n            marker_opacity=0\n        )\n    )\n\n    # configure axes\n    fig.update_xaxes(\n        visible=False,\n        range=[0, img_width]\n    )\n\n    fig.update_yaxes(\n        visible=False,\n        range=[0, img_height],\n        # the scaleanchor attribute ensures that the aspect ratio stays constant\n        scaleanchor=\"x\"\n    )\n\n    # add image\n    fig.add_layout_image(dict(\n        x=0,\n        sizex=img_width,\n        y=img_height,\n        sizey=img_height,\n        xref=\"x\", yref=\"y\",\n        opacity=1.0,\n        layer=\"below\",\n        sizing=\"stretch\",\n        source=img_example\n    ))\n    \n    # add polygons\n    for annotation in annotations:\n        name = annotation[\"type\"]\n        xs, ys = get_cartesian_coords(annotation[\"coordinates\"], img_height)\n        fig.add_trace(go.Scatter(\n            x=xs, y=ys, fill=\"toself\",\n            name=name,\n            hovertemplate=\"%{name}\",\n            mode='lines'\n        ))\n\n    # configure other layout\n    fig.update_layout(\n        width=img_width * scale_factor,\n        height=img_height * scale_factor,\n        margin={\"l\": 0, \"r\": 0, \"t\": 0, \"b\": 0},\n        showlegend=False\n    )\n\n    # disable the autosize on double click because it adds unwanted margins around the image\n    # and finally show figure\n    fig.show(config={'doubleClick': 'reset'})\n","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:58.193099Z","iopub.execute_input":"2024-08-25T11:33:58.193467Z","iopub.status.idle":"2024-08-25T11:33:58.205797Z","shell.execute_reply.started":"2024-08-25T11:33:58.193433Z","shell.execute_reply":"2024-08-25T11:33:58.204962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_annotated_image(tiles_dicts[10])","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:58.206955Z","iopub.execute_input":"2024-08-25T11:33:58.207292Z","iopub.status.idle":"2024-08-25T11:33:58.552177Z","shell.execute_reply.started":"2024-08-25T11:33:58.207263Z","shell.execute_reply":"2024-08-25T11:33:58.551199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = np.zeros((512, 512), dtype=np.float32)\nfor annot in tiles_dicts[0]['annotations']:\n    cords = annot['coordinates']\n    if annot['type'] == \"blood_vessel\":\n        for cd in cords:\n            rr, cc = np.array([i[1] for i in cd]), np.asarray([i[0] for i in cd])\n            mask[rr, cc] = 1\n    \n#     if annot['type'] == \"glomerulus\":\n#         for cd in cords:\n#             rr, cc = np.array([i[1] for i in cd]), np.asarray([i[0] for i in cd])\n#             mask[rr, cc] = 2       \n            \nplt.imshow(mask)\nplt.show()\n\ncontours,_ = cv2.findContours((mask*255).astype(np.uint8), 1, 2)\nzero_img2 = np.zeros([mask.shape[0], mask.shape[1], 3], dtype=\"uint8\")\n\nfor p in contours:\n    cv2.fillPoly(zero_img2, [p], (0, 0, 255))\n    \nplt.imshow(zero_img2)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:58.553268Z","iopub.execute_input":"2024-08-25T11:33:58.553869Z","iopub.status.idle":"2024-08-25T11:33:59.079525Z","shell.execute_reply.started":"2024-08-25T11:33:58.55384Z","shell.execute_reply":"2024-08-25T11:33:59.078595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# glom mask\nmask = np.zeros((512, 512), dtype=np.float32)\nfor annot in tiles_dicts[0]['annotations']:\n    cords = annot['coordinates']\n    if annot['type'] == \"glomerulus\":\n        for cd in cords:\n            rr, cc = np.array([i[1] for i in cd]), np.asarray([i[0] for i in cd])\n            mask[rr, cc] = 1       \n            \ncontours,_ = cv2.findContours((mask*255).astype(np.uint8), 1, 2)\nzero_img = np.zeros([mask.shape[0], mask.shape[1], 3], dtype=\"uint8\")\n\nfor p in contours:\n    cv2.fillPoly(zero_img, [p], (255,0,0))\n    \nplt.imshow(zero_img)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:59.080714Z","iopub.execute_input":"2024-08-25T11:33:59.081Z","iopub.status.idle":"2024-08-25T11:33:59.340067Z","shell.execute_reply.started":"2024-08-25T11:33:59.080975Z","shell.execute_reply":"2024-08-25T11:33:59.339119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = np.zeros((512, 512), dtype=np.float32)\nmask2 = np.zeros((512, 512), dtype=np.float32)\nfor annot in tiles_dicts[0]['annotations']:\n    cords = annot['coordinates']\n    if annot['type'] == \"blood_vessel\":\n        for cd in cords:\n            rr, cc = np.array([i[1] for i in cd]), np.asarray([i[0] for i in cd])\n            mask[rr, cc] = 1\n    \n    if annot['type'] == \"glomerulus\":\n        for cd in cords:\n            rr, cc = np.array([i[1] for i in cd]), np.asarray([i[0] for i in cd])\n            mask2[rr, cc] = 1      \n            \ncontours,_ = cv2.findContours((mask*255).astype(np.uint8), 1, 2)\nzero_img2 = np.zeros([mask.shape[0], mask.shape[1], 3], dtype=\"uint8\")\nfor p in contours:\n    cv2.fillPoly(zero_img2, [p], (0, 255, 0))\ncontours,_ = cv2.findContours((mask2*255).astype(np.uint8), 1, 2)\nzero_img = np.zeros([mask2.shape[0], mask2.shape[1], 3], dtype=\"uint8\")\nfor p in contours:\n    cv2.fillPoly(zero_img, [p], (255, 0, 0))\n\nzero_img2 + zero_img\n\nplt.imshow(zero_img2 + zero_img)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:59.341391Z","iopub.execute_input":"2024-08-25T11:33:59.341756Z","iopub.status.idle":"2024-08-25T11:33:59.610149Z","shell.execute_reply.started":"2024-08-25T11:33:59.341722Z","shell.execute_reply":"2024-08-25T11:33:59.609312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# make it a function\ndef make_seg_mask(tiles_dict):\n    mask = np.zeros((512, 512), dtype=np.float32)\n    mask2 = np.zeros((512, 512), dtype=np.float32)\n\n    for annot in tiles_dict['annotations']:\n        cords = annot['coordinates']\n        if annot['type'] == \"blood_vessel\":\n            for cd in cords:\n                rr, cc = np.array([i[1] for i in cd]), np.asarray([i[0] for i in cd])\n                mask[rr, cc] = 1\n\n        if annot['type'] == \"glomerulus\":\n            for cd in cords:\n                rr, cc = np.array([i[1] for i in cd]), np.asarray([i[0] for i in cd])\n                mask2[rr, cc] = 1      \n\n    contours,_ = cv2.findContours((mask*255).astype(np.uint8), 1, 2)\n    zero_img2 = np.zeros([mask.shape[0], mask.shape[1], 3], dtype=\"uint8\")\n    for p in contours:\n        cv2.fillPoly(zero_img2, [p], (0, 255, 0))\n    contours,_ = cv2.findContours((mask2*255).astype(np.uint8), 1, 2)\n    zero_img = np.zeros([mask2.shape[0], mask2.shape[1], 3], dtype=\"uint8\")\n    for p in contours:\n        cv2.fillPoly(zero_img, [p], (255, 0, 0))\n    \n    final = zero_img2 + zero_img\n    # glom\n    final = np.all(final == [255, 0, 0], axis=-1).astype(int) + np.all(final == [0, 255, 0], axis=-1).astype(int)*2\n    return final\n","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:59.611387Z","iopub.execute_input":"2024-08-25T11:33:59.611661Z","iopub.status.idle":"2024-08-25T11:33:59.624283Z","shell.execute_reply.started":"2024-08-25T11:33:59.611637Z","shell.execute_reply":"2024-08-25T11:33:59.623389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"array = tiff.imread('/kaggle/input/hubmap-hacking-the-human-vasculature/train/0006ff2aa7cd.tif')\nimg_example = Image.fromarray(array)\nimg = np.array(img_example)\nmask = make_seg_mask(tiles_dicts[0])\nmask","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:59.625439Z","iopub.execute_input":"2024-08-25T11:33:59.625736Z","iopub.status.idle":"2024-08-25T11:33:59.70032Z","shell.execute_reply.started":"2024-08-25T11:33:59.625713Z","shell.execute_reply":"2024-08-25T11:33:59.699314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nplt.imshow(mask, cmap='gray')  # use grayscale color map\nplt.colorbar()  # optionally add a color bar to know which color corresponds to which value\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:59.701521Z","iopub.execute_input":"2024-08-25T11:33:59.702404Z","iopub.status.idle":"2024-08-25T11:33:59.974027Z","shell.execute_reply.started":"2024-08-25T11:33:59.702375Z","shell.execute_reply":"2024-08-25T11:33:59.973203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_hehe = [tiles_dicts.pop(0),tiles_dicts.pop(), tiles_dicts.pop(len(tiles_dicts)//2)]\ntest_hehe.append(tiles_dicts.pop(len(tiles_dicts)//2))\ntest_hehe.append(tiles_dicts.pop(len(tiles_dicts)//2))\ntest_hehe.append(tiles_dicts.pop(len(tiles_dicts)//2))\ntest_hehe.append(tiles_dicts.pop(len(tiles_dicts)//2))\ntest_hehe.append(tiles_dicts.pop(len(tiles_dicts)//2))","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:59.979313Z","iopub.execute_input":"2024-08-25T11:33:59.979586Z","iopub.status.idle":"2024-08-25T11:33:59.985555Z","shell.execute_reply.started":"2024-08-25T11:33:59.979562Z","shell.execute_reply":"2024-08-25T11:33:59.98457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs('train/image', exist_ok=True)\nos.makedirs('train/mask', exist_ok=True)\n\nfor i, tldc in enumerate(tqdm(tiles_dicts)):\n    array = tiff.imread(f'/kaggle/input/hubmap-hacking-the-human-vasculature/train/{tldc[\"id\"]}.tif')\n    img_example = Image.fromarray(array)\n    img = np.array(img_example)\n    mask = make_seg_mask(tldc)\n    if np.sum(mask)>0:\n        cv2.imwrite(f'train/image/{tldc[\"id\"]}.png', img)\n        cv2.imwrite(f'train/mask/{tldc[\"id\"]}_mask.png', mask)\n        \n\n","metadata":{"execution":{"iopub.status.busy":"2024-08-25T11:33:59.986737Z","iopub.execute_input":"2024-08-25T11:33:59.987078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.makedirs('test/image', exist_ok=True)\nos.makedirs('test/mask', exist_ok=True)\n\nfor i, tldc in enumerate(tqdm(test_hehe)):\n    array = tiff.imread(f'/kaggle/input/hubmap-hacking-the-human-vasculature/train/{tldc[\"id\"]}.tif')\n    img_example = Image.fromarray(array)\n    img = np.array(img_example)\n    mask = make_seg_mask(tldc)\n    if np.sum(mask)>0:\n        cv2.imwrite(f'test/image/{tldc[\"id\"]}.png', img)\n        cv2.imwrite(f'test/mask/{tldc[\"id\"]}_mask.png', mask)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## create data","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset\nimport os\nfrom PIL import Image\nimport random\n\nclass SemanticSegmentationDataset(Dataset):\n    \"\"\"Image (semantic) segmentation dataset.\"\"\"\n    \n    # fuck it, hard code it all\n    def __init__(self,root_dir, feature_extractor, train=True, transform=None):\n        \"\"\"\n        Args:\n            root_dir (string): Root directory of the dataset containing the images + annotations.\n            feature_extractor (SegFormerFeatureExtractor): feature extractor to prepare images + segmentation maps.\n            train (bool): Whether to load \"training\" or \"validation\" images + annotations.\n        \"\"\"\n        self.root_dir = root_dir\n        self.feature_extractor = feature_extractor\n        self.train = train\n        self.transform = transform\n        \n        \n        sub_path = \"training\" if self.train else \"validation\"\n        self.img_dir = f'{root_dir}/image'\n        self.ann_dir = f'{root_dir}/mask'\n        \n        \n        image_file_names = []\n        for root, dirs, files in os.walk(self.img_dir):\n          image_file_names.extend(files)\n        self.images = sorted(image_file_names)\n        \n        # read annotations\n        annotation_file_names = []\n        for root, dirs, files in os.walk(self.ann_dir):\n          annotation_file_names.extend(files)\n        self.annotations = sorted(annotation_file_names)\n\n        assert len(self.images) == len(self.annotations), \"There must be as many images as there are segmentation maps\"\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx): \n        image = Image.open(os.path.join(self.img_dir, self.images[idx]))\n        segmentation_map = Image.open(os.path.join(self.ann_dir, self.annotations[idx]))\n        \n        if self.transform:\n            seed = np.random.randint(2147483647)\n            random.seed(seed)\n            torch.manual_seed(seed)\n            image = transform(image)\n            random.seed(seed)\n            torch.manual_seed(seed)\n            segmentation_map = transform(segmentation_map)\n            \n        # randomly crop + pad both image and segmentation map to same size?\n        encoded_inputs = self.feature_extractor(image, segmentation_map, return_tensors=\"pt\")\n        \n        for k,v in encoded_inputs.items():\n          encoded_inputs[k].squeeze_() # remove batch dimension\n        \n        return encoded_inputs\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# define a transform\n\nclass RandomCropRange:\n    def __init__(self, min_size, max_size):\n        self.min_size = min_size\n        self.max_size = max_size\n\n    def __call__(self, img):\n        crop_size = torch.randint(self.min_size, self.max_size+1, (1,)).item()\n        top = int((img.size[1]-crop_size)/2)\n        left = int((img.size[0]-crop_size)/2)\n        return transforms.functional.crop(img, top, left, crop_size, crop_size)\n    \nfrom torchvision import transforms\ntransform = transforms.Compose([\n    RandomCropRange(300,512),  \n    transforms.RandomHorizontalFlip(),  # Apply horizontal flip\n    transforms.RandomVerticalFlip(),  # Apply vertical flip\n    transforms.RandomRotation(30),\n    transforms.Resize(512)\n])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import SegformerFeatureExtractor\nroot_dir = '/kaggle/working/train'\nroot_dir_test = '/kaggle/working/test'\nfeature_extractor = SegformerFeatureExtractor(reduce_labels=False)\ntrain_dataset = SemanticSegmentationDataset(root_dir=root_dir, feature_extractor=feature_extractor, transform = transform)\ntest_dataset = SemanticSegmentationDataset(root_dir=root_dir_test, feature_extractor=feature_extractor)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Number of training examples:\", len(train_dataset))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoded_inputs = train_dataset[0]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoded_inputs[\"pixel_values\"].shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nplt.imshow(np.transpose(encoded_inputs[\"pixel_values\"], (1, 2, 0)))\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoded_inputs[\"labels\"].shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoded_inputs[\"labels\"]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoded_inputs[\"labels\"].squeeze().unique()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(encoded_inputs[\"labels\"])\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load the images and the respective masks for training\n!pip install -q datasets transformers evaluate","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\nfrom pytorch_lightning.loggers import WandbLogger\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwandb_api = user_secrets.get_secret(\"api_wandb\")\n\nepochs = 15\n\nwandb.login(key=wandb_api)\nwandb.init(\n    # set the wandb project where this run will be logged\n    project=\"Segformer\",\n    config={\n    \"learning_rate\": 1e-3,\n    \"architecture\": \"Segformer\",\n    \"dataset\": \"Caps\",\n    \"epochs\": 1,\n    }\n)\n\nwandb_logger = WandbLogger()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Wrapped Model","metadata":{}},{"cell_type":"code","source":"# !conda install lightning -c conda-forge","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(pl.__version__)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom sklearn.metrics import accuracy_score\nfrom tqdm.notebook import tqdm\nfrom PIL import Image\n\n\n\nimport pytorch_lightning as pl\nfrom torch.nn.functional import threshold, normalize\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\nfrom torchvision.utils import make_grid\n\n\ncolor_map = {\n    0: (0, 0, 0),    # Black\n    1: (255, 255, 255),  # White\n    2: (255, 0, 0)   # Red\n}\n\n\nclass MyModel(pl.LightningModule):\n    def __init__(self, model, plot_every_x_steps, learning_rate=0.00006):\n        super(MyModel, self).__init__()\n        self.model = model\n        self.plot_every_x_steps = plot_every_x_steps\n        self.batch_size = 1\n        self.learning_rate = learning_rate\n        self.lr = learning_rate\n        \n        \n    def forward_pass(self, batch):\n        pixel_values = batch[\"pixel_values\"]\n        labels = batch[\"labels\"]\n        outputs = self.model(pixel_values=pixel_values, labels=labels)\n        return outputs\n\n    def training_step(self, batch, batch_idx):\n        pixel_values = batch[\"pixel_values\"]\n        labels = batch[\"labels\"]\n        outputs = self.model(pixel_values=pixel_values, labels=labels)\n        loss, logits = outputs.loss, outputs.logits\n        \n        self.last_mask = labels[-1]\n        self.last_logits = logits[-1]\n        \n        self.log('train_loss', loss)\n        \n        self.log_images()\n        \n        return loss\n\n    @torch.no_grad()\n    def validation_step(self, batch, batch_idx):\n        pixel_values = batch[\"pixel_values\"]\n        labels = batch[\"labels\"]\n        outputs = self.model(pixel_values=pixel_values, labels=labels)\n        loss, logits = outputs.loss, outputs.logits\n        self.last_mask = labels[-1]\n        self.last_logits = logits[-1]\n        self.log('val_loss', loss)\n        self.log_images()\n        return loss\n    \n    \n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(self.model.parameters(), lr=self.lr or self.learning_rate, weight_decay=1e-4)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10, eta_min=0)\n        return {\"optimizer\": optimizer, \"lr_scheduler\": scheduler, \"monitor\": \"train_loss\"}\n     \n    def log_images(self):\n        if self.global_step % self.plot_every_x_steps == 0:  # plot every x epochs\n            last_mask = self.last_mask\n            last_logits = self.last_logits\n            \n            upsampled_logits = nn.functional.interpolate(last_logits.unsqueeze(0), size=last_mask.shape[-2:], mode=\"bilinear\", align_corners=False)\n            predicted = upsampled_logits.argmax(dim=1).cpu().numpy().squeeze()\n            last_mask = last_mask.cpu().numpy()\n\n            fig, ax = plt.subplots(1, 2, figsize=(15, 15))\n            ax[0].imshow(last_mask)\n            ax[0].title.set_text('Input Images')\n            ax[1].imshow(predicted, cmap='gray')\n            ax[1].title.set_text('Upscaled Masks')\n            plt.show()\n            \n            rgb_array = np.array([[color_map[value] for value in row] for row in last_mask], dtype=np.uint8)\n            rgb_array2 = np.array([[color_map[value] for value in row] for row in predicted], dtype=np.uint8)\n            self.logger.experiment.log({\n                \"Ground Truth Masks\": wandb.Image(Image.fromarray(rgb_array)),\n                \"Predicted (upscaled) Masks\": wandb.Image(Image.fromarray(rgb_array2)),\n            })","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\n\ncheckpoint_callback = ModelCheckpoint(\n    dirpath=\"checkpoints\",\n    filename=\"best-checkpoint\",\n    save_top_k=1,\n    verbose=True,\n    monitor=\"val_loss\",  # assumes you have a validation step where you log a \"val_loss\"\n    mode=\"min\",\n    every_n_epochs=5,  # change this to save every X epochs\n)\nearly_stopping_callback = EarlyStopping(\n    monitor=\"val_loss\",  # assumes you have a validation step where you log a \"val_loss\"\n    patience=3,\n    mode='min'\n)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import SegformerForSemanticSegmentation\nimport json\nfrom huggingface_hub import cached_download, hf_hub_url\n\n\nid2label = {0:'background', 1: 'bubu', 2:'dudu'}\nlabel2id = {v: k for k, v in id2label.items()}","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 1\nfrom torch.utils.data import DataLoader\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)\nval_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\n\ntry:\n    del model_in\n    del trainer\n    del model\n    del tuner\nexcept:\n    print('wtf')\nfinally:\n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# define model\nmodel_in = SegformerForSemanticSegmentation.from_pretrained(\"nvidia/mit-b5\",\n                                                         num_labels=3,\n                                                         id2label=id2label, \n                                                         label2id=label2id\n)\n\nclass LitDataModule(pl.LightningDataModule):\n    def __init__(self, batch_size):\n        super().__init__()\n        self.save_hyperparameters()\n        self.batch_size = batch_size\n    def train_dataloader(self):\n        return DataLoader(train_dataset, batch_size=self.batch_size | self.hparams.batch_size)\n    def val_dataloader(self):\n        return DataLoader(test_dataset, batch_size=self.batch_size | self.hparams.batch_size)\n    \ntrainer = pl.Trainer(\n    accumulate_grad_batches=4,\n    gradient_clip_val=0.5,\n    logger=wandb_logger,\n    max_epochs=epochs,\n    callbacks=[pl.callbacks.StochasticWeightAveraging(swa_lrs=0.00006), checkpoint_callback, early_stopping_callback],\n)\n\n\nmodel = MyModel(model_in, plot_every_x_steps=40)\ndatamodule = LitDataModule(batch_size=1)\ntuner = pl.tuner.Tuner(trainer)\n\ntuner.scale_batch_size(model, mode=\"power\", datamodule=datamodule)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# tuner.lr_find(model,datamodule=datamodule)\n# print(model.lr)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.fit(model, datamodule=datamodule)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from datasets import load_metric\n# metric = load_metric(\"mean_iou\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# save the weights!\ntorch.save(model.model.state_dict(), 'model_weights')","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}