{"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":"markdown","source":"This notebook extends the [tutorial notebook](https://www.kaggle.com/code/jpposma/vesuvius-challenge-ink-detection-tutorial) created by the organizers, added features:\n* train and validation split using matrix and morphological operations\n* animations of label 0 and 1 subvolumes that go into the network\n* integrated Weights & Biases for experiment tracking\n* added Dice score calculation ","metadata":{"_kg_hide-input":false,"_kg_hide-output":false}},{"cell_type":"code","source":"try:\n    from celluloid import Camera\nexcept:\n    !pip install celluloid","metadata":{"_kg_hide-input":true,"_kg_hide-output":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport numpy as np\nfrom PIL import Image\nfrom typing import List, Tuple\nfrom celluloid import Camera\nfrom IPython.display import HTML, display\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\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 wandb\nfrom kaggle_secrets import UserSecretsClient\nwandb.login(key=UserSecretsClient().get_secret(\"WANDB_API_KEY\"))","metadata":{"execution":{"iopub.status.busy":"2023-04-08T19:28:53.467317Z","iopub.execute_input":"2023-04-08T19:28:53.467881Z","iopub.status.idle":"2023-04-08T19:28:54.219989Z","shell.execute_reply.started":"2023-04-08T19:28:53.467825Z","shell.execute_reply":"2023-04-08T19:28:54.219004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection\"\nTRAIN_DIR = \"/kaggle/input/vesuvius-challenge-ink-detection/train\"\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nNUM_SLICES = 65\nIMAGE_SHAPE = 8181, 6330\nDEBUG = False\n\nconfig = {\"seed\": 0,\n          \"subvolume_size\": 30,\n          \"z_start\": 27,\n          \"z_dim\": 10,\n          \"batch_size\": 32,\n          \"train_rounds\": 6,\n          \"train_steps\": 5000}\n\nEXPERIMENT_NAME = f\"3Dclf-subv_{config['subvolume_size']}-zstrt_{config['z_start']}-zdim_{config['z_dim']}-bs_{config['batch_size']}-round_{config['train_rounds']}-step_{config['train_steps']}\"","metadata":{"execution":{"iopub.status.busy":"2023-04-08T19:28:54.222101Z","iopub.execute_input":"2023-04-08T19:28:54.222953Z","iopub.status.idle":"2023-04-08T19:28:54.230173Z","shell.execute_reply.started":"2023-04-08T19:28:54.222885Z","shell.execute_reply":"2023-04-08T19:28:54.229093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed):\n    \"\"\"\n    Sets the seed of the entire notebook for reproducibility.\n    \"\"\"\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # when running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # set a fixed value for the hash seed\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    \nset_seed(config[\"seed\"])","metadata":{"execution":{"iopub.status.busy":"2023-04-08T19:28:54.232948Z","iopub.execute_input":"2023-04-08T19:28:54.234125Z","iopub.status.idle":"2023-04-08T19:28:54.241249Z","shell.execute_reply.started":"2023-04-08T19:28:54.234084Z","shell.execute_reply":"2023-04-08T19:28:54.240171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{}},{"cell_type":"code","source":"def load_png(fragment_id: str, png_name: str) -> np.ndarray:\n    fragment_dir = os.path.join(TRAIN_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 show_pngs(mask: np.ndarray, inklabels: np.ndarray, ir: np.ndarray):\n    plt.style.use(\"default\")\n    plt.rcParams[\"figure.dpi\"] = 100    \n    fig, axs = plt.subplots(nrows=1, ncols=3, figsize=(24, 12))\n    axs.flatten()\n    show_array(mask, \"mask\", axs[0])\n    show_array(inklabels, \"inklabels\", axs[1])\n    show_array(ir, \"ir\", axs[2])\n        \n    \ndef show_rect_on_inklabels(rect: dict, inklabels: np.ndarray):\n    plt.style.use(\"default\")\n    fig, ax = plt.subplots()\n    show_array(inklabels, \"inklabels\", ax, off=False)\n    patch = Rectangle(xy=(rect[\"x\"], rect[\"y\"]), width=rect[\"width\"], height=rect[\"height\"], linewidth=2, edgecolor=\"red\", facecolor=\"none\")\n    ax.add_patch(patch)\n    plt.show()\n    \n\ndef load_volume(fragment_id: str, z_start: int, z_dim: int) -> np.ndarray:\n    volume_dir = os.path.join(TRAIN_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\nclass AverageCalc:\n    \"\"\"\n    Calculates and stores the average and current value.\n    Used to update the loss.\n    \"\"\"\n    def __init__(self):\n        self.reset()\n    \n    def reset(self):\n        self.value = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n    \n    def update(self, value, size):\n        self.value = value\n        self.sum += value * size\n        self.count += size\n        self.avg = self.sum/self.count\n        \nclass DummyLogger:\n    def log(self, message):\n        pass\n    \n    def finish(self):\n        pass\n    \ndef create_logger(config: dict):\n    logger = wandb.init(project=\"Vesuvius\", name=EXPERIMENT_NAME, config=config)\n    return logger","metadata":{"execution":{"iopub.status.busy":"2023-04-08T19:28:54.244350Z","iopub.execute_input":"2023-04-08T19:28:54.244916Z","iopub.status.idle":"2023-04-08T19:28:54.262328Z","shell.execute_reply.started":"2023-04-08T19:28:54.244869Z","shell.execute_reply":"2023-04-08T19:28:54.261111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = load_png(fragment_id=\"1\", png_name=\"mask\")\ninklabels = load_png(fragment_id=\"1\", png_name=\"inklabels\")\nir = load_png(fragment_id=\"1\", png_name=\"ir\")\nvolume = load_volume(fragment_id=\"1\", z_start=config[\"z_start\"], z_dim=config[\"z_dim\"])\nshow_pngs(mask, inklabels, ir)\ndel ir\ngc.collect();","metadata":{"execution":{"iopub.status.busy":"2023-04-08T19:28:54.264025Z","iopub.execute_input":"2023-04-08T19:28:54.264488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training validation split","metadata":{}},{"cell_type":"code","source":"rect = {\"x\": 1100, \"y\": 3500, \"width\": 700, \"height\": 950}\nshow_rect_on_inklabels(rect, inklabels)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_val_masks(mask: np.ndarray, rect: dict, subvolume_size: int) -> Tuple[np.ndarray]:\n    # erode mask so that subvolumes will be fully within the mask\n    eroded_mask = binary_erosion(mask, structure=np.ones((subvolume_size+10, subvolume_size+10)))\n    # binary mask of the rectangle\n    rect_mask = np.zeros((mask.shape), dtype=np.uint8)\n    rect_mask[rect[\"y\"] : rect[\"y\"] + rect[\"height\"], rect[\"x\"] : rect[\"x\"] + rect[\"width\"]] = 1\n    # validation set contains pixels inside the rectangle\n    val_mask = eroded_mask * rect_mask\n    # dilate rectangle mask so that training subvolumes will have no overlap with rectangle\n    dilated_rect_mask = binary_dilation(rect_mask, structure=np.ones((subvolume_size, subvolume_size)))\n    train_mask = eroded_mask * (1 - dilated_rect_mask) \n    return train_mask, val_mask        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_mask, val_mask = get_train_val_masks(mask, rect, config[\"subvolume_size\"])\nfig, axs = plt.subplots(nrows=1, ncols=2, figsize=(24, 12))\naxs.flatten()\nshow_array(train_mask, \"train_mask\", axs[0])\nshow_array(val_mask, \"val_mask\", axs[1])\n# TODO: class distribution in train and val masks","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset and DataLoaders","metadata":{}},{"cell_type":"code","source":"train_pixels = list(zip(*np.where(train_mask == 1)))\nval_pixels = list(zip(*np.where(val_mask == 1)))\ndel train_mask, val_mask\ngc.collect();","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SubvolumeDataset(Dataset):\n    def __init__(self, volume: np.ndarray, inklabels: np.ndarray, pixels: List[Tuple], subvolume_size: int):\n        self.volume = volume\n        self.inklabels = inklabels\n        # pixels in train or validation 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 = subvolume[np.newaxis, ...]\n        subvolume = torch.from_numpy(subvolume).to(torch.float32)\n        inklabel = torch.tensor(self.inklabels[y, x], dtype=torch.float32)\n        return subvolume, inklabel\n    \n    \ndef create_data_loaders(batch_size: int, train_pixels: List[Tuple], val_pixels: List[Tuple], dataset_kwargs: dict) -> Tuple[DataLoader]:\n    train_ds = SubvolumeDataset(pixels=train_pixels, **dataset_kwargs)\n    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=os.cpu_count())\n    \n    val_ds = SubvolumeDataset(pixels=val_pixels, **dataset_kwargs)\n    val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=os.cpu_count())\n    return train_loader, val_loader\n\n\ndef show_single_subvolume_batch(subvolume: Tensor, inklabel: Tensor) -> None:\n    subvolume = torch.squeeze(subvolume)\n    plt.rcParams[\"figure.dpi\"] = 350\n    plt.style.use(\"dark_background\")\n    num_slices = subvolume.shape[0]\n    fig = plt.figure()\n    camera = Camera(fig)\n    for idx in range(num_slices):\n        plt.imshow(subvolume[idx, :, :], cmap=\"gray\")\n        plt.text(x=0, y=-1, s=f\"slice {idx+1}/{num_slices}\", horizontalalignment=\"center\")\n        plt.text(x=-6, y=-3, s=f\"ink label: {inklabel.item()}\")\n        plt.axis(\"off\")\n        camera.snap()\n    animation = camera.animate()\n    plt.close(fig)\n    fix_video_adjust = \"<style> video {margin: 0px; padding: 0px; width:100%; height:auto;} </style>\"\n    display(HTML(fix_video_adjust + animation.to_html5_video()))\n    del animation\n    gc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize batches ","metadata":{}},{"cell_type":"code","source":"def find_batches_both_classes(dataloader: DataLoader) -> Tuple[Tuple]:\n    batch_inklabel_0 = None\n    batch_inklabel_1 = None\n    for subvolume, inklabel in train_loader_check:\n        if inklabel.item() == 0 and batch_inklabel_0 is None:\n            batch_inklabel_0 = subvolume, inklabel\n        if inklabel.item() == 1 and batch_inklabel_1 is None:\n            batch_inklabel_1 = subvolume, inklabel\n        if batch_inklabel_0 is not None and batch_inklabel_1 is not None:\n            break\n    return batch_inklabel_0, batch_inklabel_1","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_kwargs = {\"volume\": volume, \"inklabels\": inklabels, \"subvolume_size\": config[\"subvolume_size\"]}\nif not DEBUG:\n    train_loader_check, val_loader_check = create_data_loaders(1, train_pixels, val_pixels, dataset_kwargs)\n    batch_inklabel_0, batch_inklabel_1 = find_batches_both_classes(train_loader_check)\n    show_single_subvolume_batch(batch_inklabel_0[0], batch_inklabel_0[1])\n    del train_loader_check, val_loader_check, batch_inklabel_0\n    gc.collect();","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not DEBUG:\n    show_single_subvolume_batch(batch_inklabel_1[0], batch_inklabel_1[1])\n    del batch_inklabel_1\n    gc.collect();","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"def dice_coef_torch(preds: Tensor, targets: Tensor, beta=0.5, smooth=1e-5) -> float:\n    \"\"\"\n    https://www.kaggle.com/competitions/vesuvius-challenge-ink-detection/discussion/397288\n    \"\"\"\n    #comment out if your model contains a sigmoid or equivalent activation layer\n    #preds = torch.sigmoid(preds)\n\n    # flatten label and prediction tensors\n    preds = preds.view(-1).float()\n    targets = targets.view(-1).float()\n\n    y_true_count = targets.sum()\n    ctp = preds[targets==1].sum()\n    cfp = preds[targets==0].sum()\n    beta_squared = beta * beta\n\n    c_precision = ctp / (ctp + cfp + smooth)\n    c_recall = ctp / (y_true_count + smooth)\n    dice = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall + smooth)\n    return dice\n\n\nmodel = nn.Sequential(\n    nn.Conv3d(1, 16, 3, 1, 1), nn.MaxPool3d(2, 2),\n    nn.Conv3d(16, 32, 3, 1, 1), nn.MaxPool3d(2, 2),\n    nn.Conv3d(32, 64, 3, 1, 1), nn.MaxPool3d(2, 2),\n    nn.Flatten(start_dim=1),\n    nn.LazyLinear(128), nn.ReLU(),\n    nn.LazyLinear(1), nn.Sigmoid()\n).to(DEVICE)\n\n\ndef train_fn(train_steps: int, train_loader: DataLoader, model: nn.Module, loss_fn: nn.Module, optimizer, logger):\n    model.train()\n    for i, (x, y) in enumerate(train_loader):\n        x = x.to(DEVICE)\n        y = y.to(DEVICE)\n        \n        if DEBUG and (i > 5):\n            break\n        if i > train_steps:\n            break\n        \n        optimizer.zero_grad()\n        yhat = model(x).squeeze(1)\n        loss = loss_fn(yhat, y)\n        loss.backward()\n        optimizer.step()\n        logger.log({\"train_loss\": loss.item(),\n                    \"train_dice\": dice_coef_torch(yhat, y)})\n\n\n@torch.no_grad()\ndef val_fn(val_loader: DataLoader, model: nn.Module, loss_fn: nn.Module) -> Tuple[float, float]:\n    run_loss = AverageCalc()\n    model.eval()\n    y_val = []\n    yhat_val = []\n    \n    for i, (x, y) in enumerate(val_loader):\n        x = x.to(DEVICE)\n        y = y.to(DEVICE)\n        batch_size = x.shape[0]\n        \n        if DEBUG and (i > 5):\n            break\n\n        yhat = model(x).squeeze(1)\n        loss = loss_fn(yhat, y)\n        run_loss.update(loss.item(), x.shape[0])\n\n        y_val.append(y.cpu())\n        yhat_val.append(yhat.cpu())\n        \n    y_val = torch.cat(y_val)\n    yhat_val = torch.cat(yhat_val)\n    return run_loss.avg, dice_coef_torch(yhat_val, y_val)\n\n\n@torch.no_grad()\ndef generate_ink_pred(val_loader: DataLoader, model: nn.Module, val_pixels: List[Tuple[int, int]]) -> Tensor:\n    model.eval()\n    output = torch.zeros(IMAGE_SHAPE, dtype=torch.float32)\n    for i, (x, y) in enumerate(val_loader):\n        x = x.to(DEVICE)\n        y = y.to(DEVICE)\n        batch_size = x.shape[0]\n        \n        if DEBUG and (i > 5):\n            break\n\n        yhat = model(x).squeeze(1)\n        for j, pred in enumerate(yhat):\n            output[val_pixels[i * batch_size + j]] = pred\n    return output","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_fn = nn.BCELoss()\nlogger = DummyLogger() if DEBUG else create_logger(config)\ntrain_loader, val_loader = create_data_loaders(config[\"batch_size\"], train_pixels, val_pixels, dataset_kwargs)\ndel train_pixels\ngc.collect()\n\noptimizer = Adam(model.parameters())\nfor _ in range(config[\"train_rounds\"]):\n    train_fn(config[\"train_steps\"], train_loader, model, loss_fn, optimizer, logger)\n    val_loss, val_dice = val_fn(val_loader, model, loss_fn)\n    logger.log({\"val_loss\": val_loss, \"val_dice\": val_dice})\nlogger.finish()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ink_pred = generate_ink_pred(val_loader, model, val_pixels)\ndel val_pixels\ngc.collect()\n\nfig, axs = plt.subplots(nrows=1, ncols=2, figsize=(10, 10))\naxs.flatten()\nshow_array(ink_pred.cpu().numpy(), \"ink_pred\", axs[0])\nshow_array(inklabels, \"inklabels\", axs[1])","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}