{"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":"import os\nimport gc\nimport yaml\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom typing import List, Tuple\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\nfrom sklearn.metrics import roc_auc_score\nfrom scipy.ndimage.morphology import binary_erosion, binary_dilation\n\nimport torch\nfrom torch import nn\nfrom torch import Tensor\nfrom torch.optim import Adam\nfrom torch.utils.data import Dataset, DataLoader\n\nimport sys\nsys.path.append(\"../input/timm-pytorch-image-models/pytorch-image-models-master\")\nimport timm","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:46:50.376747Z","iopub.execute_input":"2023-04-21T07:46:50.377162Z","iopub.status.idle":"2023-04-21T07:46:50.384684Z","shell.execute_reply.started":"2023-04-21T07:46:50.377125Z","shell.execute_reply":"2023-04-21T07:46:50.383165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nTEST_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection/test\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nDEBUG = False\n\n# TO UPDATE\nCONFIG_PATH = \"/kaggle/input/train-2-5d-pretrained-effnet-b0-ink-classifier/wandb/run-20230415_091904-9lv6k4im/files/config.yaml\"\nMODEL_PATH = \"/kaggle/input/train-2-5d-pretrained-effnet-b0-ink-classifier/efficientnet_b0-subv_30-zstrt_27-zdim_10-bs_32-round_3-step_5000-thr_0.6.pt\"","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:46:50.391546Z","iopub.execute_input":"2023-04-21T07:46:50.391877Z","iopub.status.idle":"2023-04-21T07:46:50.397565Z","shell.execute_reply.started":"2023-04-21T07:46:50.391847Z","shell.execute_reply":"2023-04-21T07:46:50.396070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"def load_config(path: str) -> dict:\n    with open(path, \"r\") as file:\n        config = yaml.safe_load(file)\n    return config\n\ndef convert_config(config: dict) -> dict:\n    keys_to_remove = [\"wandb_version\", \"_wandb\"]\n    filtered_config = {key: value for key, value in config.items() if key not in keys_to_remove}\n    new_config = {}\n    for key, value in filtered_config.items():\n        new_config[key] = value[\"value\"]\n    return new_config\n\n\ndef load_png(fragment_id: str, png_name: str) -> np.ndarray:\n    fragment_dir = os.path.join(TEST_DIR, fragment_id)\n    path = os.path.join(fragment_dir, f\"{png_name}.png\")\n    image = Image.open(path)\n    return np.array(image)\n\n\ndef show_array(array: np.ndarray, name: str, ax, off=True):\n    ax.imshow(array, cmap=\"gray\")\n    if off:\n        ax.axis(\"off\")\n    ax.set_title(f\"{name}, shape: {array.shape}\")\n    \n\ndef load_volume(fragment_id: str, z_start: int, z_dim: int) -> np.ndarray:\n    volume_dir = os.path.join(TEST_DIR, fragment_id, \"surface_volume\")\n    volume = []\n    for i in range(z_start, z_start + z_dim):\n        slice_path = os.path.join(volume_dir, f\"{i:02d}.tif\")\n        slice_png = Image.open(slice_path)\n        # normalize pixel intesity values into [0,1]\n        slice_array = np.array(slice_png, dtype=np.float32) / 65535.0\n        volume.append(slice_array)\n    return np.stack(volume, axis=0)\n\n\ndef erode_mask(mask: np.ndarray, subvolume_size: int):\n    eroded_mask = binary_erosion(mask, structure=np.ones((2*subvolume_size, 2*subvolume_size)))\n    return eroded_mask\n\n\ndef pad_array(array: np.ndarray, subvolume_size: int) -> np.ndarray:\n    padded_array = np.pad(array, pad_width=subvolume_size+1, mode=\"constant\")\n    return padded_array","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:46:50.404738Z","iopub.execute_input":"2023-04-21T07:46:50.405019Z","iopub.status.idle":"2023-04-21T07:46:50.418116Z","shell.execute_reply.started":"2023-04-21T07:46:50.404992Z","shell.execute_reply":"2023-04-21T07:46:50.416822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb_config = load_config(CONFIG_PATH)\nconfig = convert_config(wandb_config)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:46:50.420544Z","iopub.execute_input":"2023-04-21T07:46:50.422049Z","iopub.status.idle":"2023-04-21T07:46:50.446550Z","shell.execute_reply.started":"2023-04-21T07:46:50.422006Z","shell.execute_reply":"2023-04-21T07:46:50.445512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask_a = load_png(fragment_id=\"a\", png_name=\"mask\")\nmask_a = erode_mask(mask_a, config[\"subvolume_size\"])\n\nmask_b = load_png(fragment_id=\"b\", png_name=\"mask\")\nmask_b = erode_mask(mask_b, config[\"subvolume_size\"])\n\nfig, axs = plt.subplots(nrows=1, ncols=2, figsize=(12, 12))\naxs.flatten()\nshow_array(mask_a, \"mask_a\", axs[0])\nshow_array(mask_b, \"mask_b\", axs[1])","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:46:50.448170Z","iopub.execute_input":"2023-04-21T07:46:50.448526Z","iopub.status.idle":"2023-04-21T07:48:14.798011Z","shell.execute_reply.started":"2023-04-21T07:46:50.448490Z","shell.execute_reply":"2023-04-21T07:48:14.796807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SubvolumeDatasetTest(Dataset):\n    def __init__(self, volume: np.ndarray, pixels: List[Tuple], subvolume_size: int):\n        self.volume = volume\n        # pixels in test mask\n        self.pixels = pixels\n        self.subvolume_size = subvolume_size\n        \n    def __len__(self):\n        return len(self.pixels)\n    \n    def __getitem__(self, idx):\n        y, x = self.pixels[idx]\n        subvolume = self.volume[:, \n                                y - self.subvolume_size : y + self.subvolume_size, \n                                x - self.subvolume_size : x + self.subvolume_size]\n        subvolume = torch.from_numpy(subvolume).to(torch.float32)\n        return subvolume\n    \n\ndef create_data_loader(batch_size: int, volume: np.ndarray, pixels: List[Tuple], subvolume_size: int):\n    dataset = SubvolumeDatasetTest(volume, pixels, subvolume_size)\n    data_loader = DataLoader(dataset, batch_size=batch_size, shuffle=False)\n    return data_loader","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:23:40.242498Z","iopub.execute_input":"2023-04-21T07:23:40.243059Z","iopub.status.idle":"2023-04-21T07:23:40.258727Z","shell.execute_reply.started":"2023-04-21T07:23:40.243023Z","shell.execute_reply":"2023-04-21T07:23:40.251259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class InkClassifier(nn.Module):\n    \n    def __init__(self, config):\n        super().__init__()\n        self.backbone = timm.create_model(config[\"model_name\"], pretrained=False, in_chans=10, num_classes=0)\n        self.backbone_dim = self.backbone(torch.rand(1, config[\"z_dim\"], 2*config[\"subvolume_size\"], 2*config[\"subvolume_size\"])).shape[-1]\n        self.classifier = nn.Linear(in_features=self.backbone_dim, out_features=1)\n        self.sigmoid = nn.Sigmoid()\n        \n    def forward(self, x):\n        x = self.backbone(x)\n        logits = self.classifier(x)\n        out = self.sigmoid(logits)\n        return out.flatten()\n    \n    \ndef load_model(path: str, config: dict):\n    model = InkClassifier(config)\n    model.load_state_dict(torch.load(path))\n    model.to(DEVICE)\n    return model\n\n\n@torch.no_grad()\ndef infer_fn(test_loader: DataLoader, model: nn.Module, image_shape_2d: Tuple[int, int], pixels: List[Tuple[int, int]]) -> torch.tensor:\n    model.eval()\n    output = torch.zeros(image_shape_2d, dtype=torch.float32)\n    for i, x in enumerate(test_loader):\n        x = x.to(DEVICE)\n        batch_size = x.shape[0]\n        \n        if DEBUG and (i > 5):\n            break\n\n        yhat = model(x)\n        for j, pred in enumerate(yhat):\n            output[pixels[i * batch_size + j]] = pred\n    return output\n\n\ndef rle(image):\n    \"\"\"\n    Convert a binary image to run-length encoding format.\n\n    Args:\n        image (numpy.ndarray): A 2D binary image.\n\n    Returns:\n        str: The RLE-encoded representation of the input image.\n    \"\"\"\n    pixels = image.flatten()\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       \n    \ndef infer_for_fragment(fragment_id: str, config: dict, model: nn.Module) -> np.ndarray:\n    mask = load_png(fragment_id, png_name=\"mask\")\n    mask = erode_mask(mask, config[\"subvolume_size\"])\n    image_shape_2d = mask.shape\n    pixels = list(zip(*np.where(mask == 1)))\n    del mask\n    gc.collect()\n    \n    volume = load_volume(fragment_id, z_start=config[\"z_start\"], z_dim=config[\"z_dim\"])\n    test_loader = create_data_loader(config[\"batch_size\"], volume, pixels, config[\"subvolume_size\"])\n    ink_pred = infer_fn(test_loader, model, image_shape_2d, pixels)\n    return ink_pred.gt(config[\"threshold\"]).cpu().numpy().astype(int)","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:23:40.260818Z","iopub.execute_input":"2023-04-21T07:23:40.261453Z","iopub.status.idle":"2023-04-21T07:23:40.294733Z","shell.execute_reply.started":"2023-04-21T07:23:40.261337Z","shell.execute_reply":"2023-04-21T07:23:40.293563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = load_model(MODEL_PATH, config)\nfragment_ids = os.listdir(TEST_DIR)\nink_preds_dict = {}\nsubmission = []\nfor fragment_id in fragment_ids:\n    ink_pred = infer_for_fragment(fragment_id, config, model)\n    ink_preds_dict[fragment_id] = ink_pred\n    submission.append({\"Id\": fragment_id, \"Predicted\": rle(ink_pred)})\nsubmission = pd.DataFrame(submission)\nnp.save(\"ink_pred_a.npy\", ink_preds_dict[\"a\"])\nnp.save(\"ink_pred_b.npy\", ink_preds_dict[\"b\"])","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:52:08.318914Z","iopub.execute_input":"2023-04-21T07:52:08.319322Z","iopub.status.idle":"2023-04-21T07:52:59.263415Z","shell.execute_reply.started":"2023-04-21T07:52:08.319284Z","shell.execute_reply":"2023-04-21T07:52:59.261930Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axs = plt.subplots(nrows=1, ncols=2, figsize=(12, 12))\naxs.flatten()\nshow_array(ink_preds_dict[\"a\"], \"ink_pred_a\", axs[0])\nshow_array(ink_preds_dict[\"b\"], \"ink_pred_b\", axs[1])","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:50:24.883305Z","iopub.execute_input":"2023-04-21T07:50:24.883844Z","iopub.status.idle":"2023-04-21T07:50:29.844769Z","shell.execute_reply.started":"2023-04-21T07:50:24.883801Z","shell.execute_reply":"2023-04-21T07:50:29.843637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv(\"submission.csv\", index=False)\nsubmission","metadata":{"execution":{"iopub.status.busy":"2023-04-21T07:29:40.194458Z","iopub.execute_input":"2023-04-21T07:29:40.194990Z","iopub.status.idle":"2023-04-21T07:29:40.211523Z","shell.execute_reply.started":"2023-04-21T07:29:40.194942Z","shell.execute_reply":"2023-04-21T07:29:40.210598Z"},"trusted":true},"execution_count":null,"outputs":[]}]}