{"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":"try:\n    from celluloid import Camera\nexcept:\n    !pip install celluloid","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"papermill":{"duration":11.538106,"end_time":"2023-04-11T18:07:50.680012","exception":false,"start_time":"2023-04-11T18:07:39.141906","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-14T21:39:43.771751Z","iopub.execute_input":"2023-04-14T21:39:43.772311Z","iopub.status.idle":"2023-04-14T21:39:55.960268Z","shell.execute_reply.started":"2023-04-14T21:39:43.772268Z","shell.execute_reply":"2023-04-14T21:39:55.958825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport numpy as np\nimport pandas as pd\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 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 wandb\nfrom kaggle_secrets import UserSecretsClient\nwandb.login(key=UserSecretsClient().get_secret(\"WANDB_API_KEY\"))\n\nimport sys\nsys.path.append(\"../input/timm-pytorch-image-models/pytorch-image-models-master\")\nimport timm","metadata":{"papermill":{"duration":7.517554,"end_time":"2023-04-11T18:07:58.204103","exception":false,"start_time":"2023-04-11T18:07:50.686549","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-14T21:39:55.962780Z","iopub.execute_input":"2023-04-14T21:39:55.963088Z","iopub.status.idle":"2023-04-14T21:40:02.439693Z","shell.execute_reply.started":"2023-04-14T21:39:55.963055Z","shell.execute_reply":"2023-04-14T21:40:02.438660Z"},"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\nLOG_TABLE = True\nDEBUG = False\n\nconfig = {\"model_name\": \"efficientnet_b0\",\n          \"subvolume_size\": 30,\n          \"z_start\": 27,\n          \"z_dim\": 10,\n          \"batch_size\": 32,\n          \"train_rounds\": 3,\n          \"train_steps\": 5000,\n          \"threshold\": 0.6,\n          \"seed\": 0}\n\nEXPERIMENT_NAME = f\"{config['model_name']}-subv_{config['subvolume_size']}-zstrt_{config['z_start']}-zdim_{config['z_dim']}-bs_{config['batch_size']}-round_{config['train_rounds']}-step_{config['train_steps']}-thr_{config['threshold']}\"","metadata":{"papermill":{"duration":0.076932,"end_time":"2023-04-11T18:07:58.287199","exception":false,"start_time":"2023-04-11T18:07:58.210267","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-14T21:40:02.441296Z","iopub.execute_input":"2023-04-14T21:40:02.441653Z","iopub.status.idle":"2023-04-14T21:40:02.514314Z","shell.execute_reply.started":"2023-04-14T21:40:02.441616Z","shell.execute_reply":"2023-04-14T21:40:02.513214Z"},"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":{"papermill":{"duration":0.016895,"end_time":"2023-04-11T18:07:58.309905","exception":false,"start_time":"2023-04-11T18:07:58.293010","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-14T21:40:02.517128Z","iopub.execute_input":"2023-04-14T21:40:02.517755Z","iopub.status.idle":"2023-04-14T21:40:02.526791Z","shell.execute_reply.started":"2023-04-14T21:40:02.517716Z","shell.execute_reply":"2023-04-14T21:40:02.525815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{"papermill":{"duration":0.005669,"end_time":"2023-04-11T18:07:58.321160","exception":false,"start_time":"2023-04-11T18:07:58.315491","status":"completed"},"tags":[]}},{"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        \n        \nclass DummyLogger:\n    def log(self, message):\n        pass\n    \n    def finish(self):\n        pass\n    \n    \ndef create_logger(config: dict):\n    logger = wandb.init(project=\"Vesuvius\", name=EXPERIMENT_NAME, config=config)\n    return logger","metadata":{"papermill":{"duration":0.024824,"end_time":"2023-04-11T18:07:58.351588","exception":false,"start_time":"2023-04-11T18:07:58.326764","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-14T21:40:02.528428Z","iopub.execute_input":"2023-04-14T21:40:02.528994Z","iopub.status.idle":"2023-04-14T21:40:02.545747Z","shell.execute_reply.started":"2023-04-14T21:40:02.528956Z","shell.execute_reply":"2023-04-14T21:40:02.544963Z"},"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":{"papermill":{"duration":32.264848,"end_time":"2023-04-11T18:08:30.622445","exception":false,"start_time":"2023-04-11T18:07:58.357597","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-14T21:40:02.547279Z","iopub.execute_input":"2023-04-14T21:40:02.547977Z","iopub.status.idle":"2023-04-14T21:40:31.351319Z","shell.execute_reply.started":"2023-04-14T21:40:02.547932Z","shell.execute_reply":"2023-04-14T21:40:31.350430Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training validation split","metadata":{"papermill":{"duration":0.012157,"end_time":"2023-04-11T18:08:30.647349","exception":false,"start_time":"2023-04-11T18:08:30.635192","status":"completed"},"tags":[]}},{"cell_type":"code","source":"rect = {\"x\": 1100, \"y\": 3400, \"width\": 1600, \"height\": 1300}\nshow_rect_on_inklabels(rect, inklabels)","metadata":{"papermill":{"duration":1.839995,"end_time":"2023-04-11T18:08:32.501268","exception":false,"start_time":"2023-04-11T18:08:30.661273","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-14T21:42:03.865727Z","iopub.execute_input":"2023-04-14T21:42:03.866106Z","iopub.status.idle":"2023-04-14T21:42:05.636300Z","shell.execute_reply.started":"2023-04-14T21:42:03.866072Z","shell.execute_reply":"2023-04-14T21:42:05.634413Z"},"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":{"execution":{"iopub.execute_input":"2023-04-11T18:08:32.530890Z","iopub.status.busy":"2023-04-11T18:08:32.529933Z","iopub.status.idle":"2023-04-11T18:08:32.537396Z","shell.execute_reply":"2023-04-11T18:08:32.536302Z"},"papermill":{"duration":0.025319,"end_time":"2023-04-11T18:08:32.539799","exception":false,"start_time":"2023-04-11T18:08:32.514480","status":"completed"},"tags":[]},"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])","metadata":{"execution":{"iopub.execute_input":"2023-04-11T18:08:32.568085Z","iopub.status.busy":"2023-04-11T18:08:32.566453Z","iopub.status.idle":"2023-04-11T18:10:01.204126Z","shell.execute_reply":"2023-04-11T18:10:01.202970Z"},"papermill":{"duration":88.667323,"end_time":"2023-04-11T18:10:01.219816","exception":false,"start_time":"2023-04-11T18:08:32.552493","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"ink class in train mask: {np.mean(inklabels[train_mask == 1]):.4f}\")\nprint(f\"ink class in val mask: {np.mean(inklabels[val_mask == 1]):.4f}\") ","metadata":{"execution":{"iopub.execute_input":"2023-04-11T18:10:01.251055Z","iopub.status.busy":"2023-04-11T18:10:01.249297Z","iopub.status.idle":"2023-04-11T18:10:01.379854Z","shell.execute_reply":"2023-04-11T18:10:01.378441Z"},"papermill":{"duration":0.149081,"end_time":"2023-04-11T18:10:01.382842","exception":false,"start_time":"2023-04-11T18:10:01.233761","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset and DataLoaders","metadata":{"papermill":{"duration":0.013214,"end_time":"2023-04-11T18:10:01.409703","exception":false,"start_time":"2023-04-11T18:10:01.396489","status":"completed"},"tags":[]}},{"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":{"execution":{"iopub.execute_input":"2023-04-11T18:10:01.437931Z","iopub.status.busy":"2023-04-11T18:10:01.437540Z","iopub.status.idle":"2023-04-11T18:10:08.490770Z","shell.execute_reply":"2023-04-11T18:10:08.489733Z"},"papermill":{"duration":7.070428,"end_time":"2023-04-11T18:10:08.493301","exception":false,"start_time":"2023-04-11T18:10:01.422873","status":"completed"},"tags":[]},"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 = 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_loader(batch_size: int, pixels: List[Tuple], dataset_kwargs: dict, shuffle: bool) -> DataLoader:\n    dataset = SubvolumeDataset(pixels=pixels, **dataset_kwargs)\n    data_loader = DataLoader(dataset, batch_size=batch_size, shuffle=shuffle)\n    return data_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":{"execution":{"iopub.execute_input":"2023-04-11T18:10:08.522376Z","iopub.status.busy":"2023-04-11T18:10:08.521852Z","iopub.status.idle":"2023-04-11T18:10:08.534603Z","shell.execute_reply":"2023-04-11T18:10:08.533358Z"},"papermill":{"duration":0.030222,"end_time":"2023-04-11T18:10:08.537495","exception":false,"start_time":"2023-04-11T18:10:08.507273","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize batches ","metadata":{"papermill":{"duration":0.013318,"end_time":"2023-04-11T18:10:08.564818","exception":false,"start_time":"2023-04-11T18:10:08.551500","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def find_batches_both_classes(data_loader: DataLoader) -> Tuple[Tuple]:\n    batch_inklabel_0 = None\n    batch_inklabel_1 = None\n    for subvolume, inklabel in data_loader:\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":{"execution":{"iopub.execute_input":"2023-04-11T18:10:08.593067Z","iopub.status.busy":"2023-04-11T18:10:08.592744Z","iopub.status.idle":"2023-04-11T18:10:08.599211Z","shell.execute_reply":"2023-04-11T18:10:08.597991Z"},"papermill":{"duration":0.023151,"end_time":"2023-04-11T18:10:08.601398","exception":false,"start_time":"2023-04-11T18:10:08.578247","status":"completed"},"tags":[]},"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 = create_data_loader(1, train_pixels, dataset_kwargs, shuffle=True)\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, batch_inklabel_0\n    gc.collect();","metadata":{"execution":{"iopub.execute_input":"2023-04-11T18:10:08.629922Z","iopub.status.busy":"2023-04-11T18:10:08.629042Z","iopub.status.idle":"2023-04-11T18:10:16.865516Z","shell.execute_reply":"2023-04-11T18:10:16.864310Z"},"papermill":{"duration":8.253838,"end_time":"2023-04-11T18:10:16.868429","exception":false,"start_time":"2023-04-11T18:10:08.614591","status":"completed"},"tags":[]},"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":{"execution":{"iopub.execute_input":"2023-04-11T18:10:16.908579Z","iopub.status.busy":"2023-04-11T18:10:16.908268Z","iopub.status.idle":"2023-04-11T18:10:21.818558Z","shell.execute_reply":"2023-04-11T18:10:21.817331Z"},"papermill":{"duration":4.932948,"end_time":"2023-04-11T18:10:21.821442","exception":false,"start_time":"2023-04-11T18:10:16.888494","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{"papermill":{"duration":0.024771,"end_time":"2023-04-11T18:10:21.872745","exception":false,"start_time":"2023-04-11T18:10:21.847974","status":"completed"},"tags":[]}},{"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=True, 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()","metadata":{"execution":{"iopub.execute_input":"2023-04-11T18:10:21.933916Z","iopub.status.busy":"2023-04-11T18:10:21.933521Z","iopub.status.idle":"2023-04-11T18:10:21.941764Z","shell.execute_reply":"2023-04-11T18:10:21.940590Z"},"papermill":{"duration":0.045686,"end_time":"2023-04-11T18:10:21.944714","exception":false,"start_time":"2023-04-11T18:10:21.899028","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"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\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)\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        try:\n            train_auc = roc_auc_score(y.detach().cpu().numpy(), yhat.detach().cpu().numpy())\n            logger.log({\"train_auc\": train_auc})\n        except ValueError:\n            pass\n\n\n@torch.no_grad()\ndef val_fn(val_loader: DataLoader, model: nn.Module, loss_fn: nn.Module) -> dict:\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)\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    df_val = pd.DataFrame({\"y_val\": y_val, \"yhat_val\": yhat_val})\n    \n    try:\n        val_auc = roc_auc_score(y_val.numpy(), yhat_val.numpy())\n    except ValueError:\n        val_auc = np.nan\n        \n    return {\"val_loss\": run_loss.avg, \"val_dice\": dice_coef_torch(yhat_val, y_val), \n            \"val_auc\": val_auc, \"df_val\": df_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)\n        for j, pred in enumerate(yhat):\n            output[val_pixels[i * batch_size + j]] = pred\n    return output","metadata":{"execution":{"iopub.execute_input":"2023-04-11T18:10:21.998928Z","iopub.status.busy":"2023-04-11T18:10:21.998625Z","iopub.status.idle":"2023-04-11T18:10:22.017819Z","shell.execute_reply":"2023-04-11T18:10:22.017138Z"},"papermill":{"duration":0.048983,"end_time":"2023-04-11T18:10:22.020000","exception":false,"start_time":"2023-04-11T18:10:21.971017","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = InkClassifier(config).to(DEVICE)\nloss_fn = nn.BCELoss()\nlogger = DummyLogger() if DEBUG else create_logger(config)\n\n\ntrain_loader = create_data_loader(config[\"batch_size\"], train_pixels, dataset_kwargs, shuffle=True)\nval_loader = create_data_loader(config[\"batch_size\"], val_pixels, dataset_kwargs, shuffle=False)\noptimizer = Adam(model.parameters())\nbest_val_dice = 0\nfor rnd in range(config[\"train_rounds\"]):\n    train_fn(config[\"train_steps\"], train_loader, model, loss_fn, optimizer, logger)\n    val_dict = val_fn(val_loader, model, loss_fn)\n    logger.log({\"val_loss\": val_dict[\"val_loss\"], \"val_dice\": val_dict[\"val_dice\"], \"val_auc\": val_dict[\"val_auc\"]})\n    if val_dict[\"val_dice\"] > best_val_dice:\n        best_val_dice = val_dict[\"val_dice\"]\n        torch.save(model.state_dict(), f\"{EXPERIMENT_NAME}.pt\")\n        print(f\"Model saved at round {rnd} with val_dice {val_dict['val_dice']:.4f}.\")\n        if LOG_TABLE:\n            logger.log({\"table\": wandb.Table(dataframe=val_dict[\"df_val\"])})\nlogger.finish()","metadata":{"execution":{"iopub.execute_input":"2023-04-11T18:10:22.075469Z","iopub.status.busy":"2023-04-11T18:10:22.074840Z","iopub.status.idle":"2023-04-11T19:13:01.032971Z","shell.execute_reply":"2023-04-11T19:13:01.031995Z"},"papermill":{"duration":3758.988656,"end_time":"2023-04-11T19:13:01.035275","exception":false,"start_time":"2023-04-11T18:10:22.046619","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ink_pred = generate_ink_pred(val_loader, model, val_pixels)\nif not DEBUG:\n    del train_pixels, val_pixels\n    gc.collect()\n\n    \nfig, axs = plt.subplots(nrows=1, ncols=2, figsize=(10, 10))\naxs.flatten()\nshow_array(ink_pred.gt(config[\"threshold\"]).cpu().numpy(), \"ink_pred\", axs[0])\nshow_array(inklabels, \"inklabels\", axs[1])","metadata":{"execution":{"iopub.execute_input":"2023-04-11T19:13:01.099664Z","iopub.status.busy":"2023-04-11T19:13:01.098482Z","iopub.status.idle":"2023-04-11T19:26:30.046610Z","shell.execute_reply":"2023-04-11T19:26:30.045508Z"},"papermill":{"duration":809.021972,"end_time":"2023-04-11T19:26:30.088845","exception":false,"start_time":"2023-04-11T19:13:01.066873","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}