{"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":"In this Notebook, I do some improvement for baseline written by **Simayi** here: https://www.kaggle.com/code/bibanh/pytorch-0-12-deeplabv3-bceloss\n- MANet Architecture\n- Combine (BCELoss + DiceLoss)","metadata":{"papermill":{"duration":0.006144,"end_time":"2023-03-31T17:06:16.333927","exception":false,"start_time":"2023-03-31T17:06:16.327783","status":"completed"},"tags":[]}},{"cell_type":"code","source":"is_train = True\nis_infer = True\nTHRESHOLD = 0.20\n#BEST_MODEL_PATH = '/kaggle/input/ink-detection-trained-model/best-manet-effnetb7.pt'","metadata":{"execution":{"iopub.status.busy":"2023-04-02T07:52:10.785811Z","iopub.execute_input":"2023-04-02T07:52:10.786598Z","iopub.status.idle":"2023-04-02T07:52:10.792362Z","shell.execute_reply.started":"2023-04-02T07:52:10.786554Z","shell.execute_reply":"2023-04-02T07:52:10.791232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Install  Segmentation-models-pytorch","metadata":{"papermill":{"duration":0.004726,"end_time":"2023-03-31T17:06:16.343861","exception":false,"start_time":"2023-03-31T17:06:16.339135","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!cp -rf /kaggle/input/segmentation-models-pytorch-032 /kaggle/working\n!pip install -q /kaggle/working/segmentation-models-pytorch-032/pretrainedmodels-0.7.4\n!pip install -q /kaggle/working/segmentation-models-pytorch-032/efficientnet_pytorch-0.7.1\n!pip install -q /kaggle/working/segmentation-models-pytorch-032/timm-0.6.12-py3-none-any.whl\n!pip install -q /kaggle/working/segmentation-models-pytorch-032/segmentation_models_pytorch-0.3.2-py3-none-any.whl","metadata":{"papermill":{"duration":122.527051,"end_time":"2023-03-31T17:08:18.875854","exception":false,"start_time":"2023-03-31T17:06:16.348803","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T07:52:10.797938Z","iopub.execute_input":"2023-04-02T07:52:10.798264Z","iopub.status.idle":"2023-04-02T07:54:15.435557Z","shell.execute_reply.started":"2023-04-02T07:52:10.798221Z","shell.execute_reply":"2023-04-02T07:54:15.434250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Copy pretrained model to local","metadata":{"papermill":{"duration":0.005489,"end_time":"2023-03-31T17:08:18.887076","exception":false,"start_time":"2023-03-31T17:08:18.881587","status":"completed"},"tags":[]}},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n!cp /kaggle/input/pretrained-pytorch/resnet101-5d3b4d8f.pth /root/.cache/torch/hub/checkpoints/\n!ls /root/.cache/torch/hub/checkpoints/","metadata":{"papermill":{"duration":6.769394,"end_time":"2023-03-31T17:08:25.661824","exception":false,"start_time":"2023-03-31T17:08:18.892430","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T07:54:15.441761Z","iopub.execute_input":"2023-04-02T07:54:15.442210Z","iopub.status.idle":"2023-04-02T07:54:22.331735Z","shell.execute_reply.started":"2023-04-02T07:54:15.442155Z","shell.execute_reply":"2023-04-02T07:54:22.330495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define Model","metadata":{"papermill":{"duration":0.005314,"end_time":"2023-03-31T17:08:25.672927","exception":false,"start_time":"2023-03-31T17:08:25.667613","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from torchvision import models\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport segmentation_models_pytorch as smp\n\ndef create_model(in_ch):\n    model = smp.Unet(\n    encoder_name=\"resnet101\",\n    encoder_weights=\"imagenet\",\n    in_channels=in_ch,\n    classes=1,\n    )\n    model.train()\n    return model","metadata":{"papermill":{"duration":4.777017,"end_time":"2023-03-31T17:08:30.455359","exception":false,"start_time":"2023-03-31T17:08:25.678342","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T07:54:22.334157Z","iopub.execute_input":"2023-04-02T07:54:22.334606Z","iopub.status.idle":"2023-04-02T07:54:26.101888Z","shell.execute_reply.started":"2023-04-02T07:54:22.334562Z","shell.execute_reply":"2023-04-02T07:54:26.100626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define Loss: BCELoss + DiceLoss","metadata":{"papermill":{"duration":0.00534,"end_time":"2023-03-31T17:08:30.466519","exception":false,"start_time":"2023-03-31T17:08:30.461179","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def dice_loss(pred, target, smooth = 1.):\n    pred = pred.contiguous()\n    target = target.contiguous()    \n\n    intersection = (pred * target).sum(dim=2).sum(dim=2)\n    \n    loss = (1 - ((2. * intersection + smooth) / (pred.sum(dim=2).sum(dim=2) + target.sum(dim=2).sum(dim=2) + smooth)))\n    \n    return loss.mean()\n\ndef calc_loss(pred, target, bce_weight = 1):\n    bce = F.binary_cross_entropy_with_logits(pred, target)\n\n    pred = F.sigmoid(pred)\n    dice = dice_loss(pred, target)\n\n    loss = bce * bce_weight + dice * (1 - bce_weight)\n\n    return loss","metadata":{"papermill":{"duration":0.015745,"end_time":"2023-03-31T17:08:30.487745","exception":false,"start_time":"2023-03-31T17:08:30.472000","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T07:54:26.105746Z","iopub.execute_input":"2023-04-02T07:54:26.106273Z","iopub.status.idle":"2023-04-02T07:54:26.114655Z","shell.execute_reply.started":"2023-04-02T07:54:26.106226Z","shell.execute_reply":"2023-04-02T07:54:26.113017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define Dataset\nHere the RandomPatchLocDataset will draw random patch from the volume.\n\n**Full visualize** can ref: https://www.kaggle.com/code/fchollet/keras-starter-kit-unet-train-on-full-dataset\n\nOn how the patch is draw / how the volume is create by concat / Where the validation set is\n\nThe notebook is the same setting with it.\n","metadata":{"papermill":{"duration":0.005302,"end_time":"2023-03-31T17:08:30.498829","exception":false,"start_time":"2023-03-31T17:08:30.493527","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import numpy as np\nimport torch.utils.data as data\nimport os\nimport PIL.Image as Image\nfrom tqdm import tqdm\nimport glob\nimport torch.nn as nn\nfrom torch import optim\nimport torch\nfrom torch.utils.tensorboard import SummaryWriter\n\n# ================== random patch dataset ==============================\nclass RandomOpt():\n    def __init__(self):\n        self.SHARED_HEIGHT = 4096  # Height to resize all papyrii\n        self.BUFFER = 64  # Half-size of papyrus patches we'll use as model inputs\n        self.Z_DIM = 16  # Number of slices in the z direction. Max value is 64 - Z_START\n        self.Z_START = 25  # Offset of slices in the z direction\n        self.DATA_DIR = \"../input/vesuvius-challenge-ink-detection\"\n\ndef resize(img, SHARED_HEIGHT=RandomOpt().SHARED_HEIGHT):\n    current_width, current_height = img.size\n    aspect_ratio = current_width / current_height\n    new_width = int(SHARED_HEIGHT * aspect_ratio)\n    new_size = (new_width, SHARED_HEIGHT)\n    img = img.resize(new_size)\n    return img\n\ndef load_mask(split, index, DATA_DIR=RandomOpt().DATA_DIR):\n    img = Image.open(f\"{DATA_DIR}/{split}/{index}/mask.png\").convert('1')\n    img = resize(img)\n    return torch.from_numpy(np.array(img))\n\ndef load_labels(split, index, DATA_DIR=RandomOpt().DATA_DIR):\n    img = Image.open(f\"{DATA_DIR}/{split}/{index}/inklabels.png\")\n    img = resize(img)\n    return torch.from_numpy(np.array(img)).gt(0).float()\n\ndef load_volume(split, index, DATA_DIR=RandomOpt().DATA_DIR, Z_START=RandomOpt().Z_START, Z_DIM=RandomOpt().Z_DIM):\n    # Load the 3d x-ray scan, one slice at a time\n    z_slices_fnames = sorted(glob.glob(f\"{DATA_DIR}/{split}/{index}/surface_volume/*.tif\"))[Z_START:Z_START + Z_DIM]\n    z_slices = []\n    for z, filename in  tqdm(enumerate(z_slices_fnames)):\n        img = Image.open(filename)\n        img = resize(img)\n        z_slice = np.array(img, dtype=\"float32\")\n        z_slices.append(torch.from_numpy(z_slice))\n    return torch.stack(z_slices, dim=0)\n\n# Random choice of patches for training\ndef sample_random_location(shape, BUFFER=RandomOpt().BUFFER):\n    a=BUFFER\n    random_train_x = (shape[0] - BUFFER - 1 - a)*torch.rand(1)+a\n    random_train_y = (shape[1] - BUFFER - 1 - a)*torch.rand(1)+a\n    random_train_location = torch.stack([random_train_x, random_train_y])\n    return random_train_location\n\ndef is_in_masked_zone(location, mask):\n    return mask[location[0].long(), location[1].long()]\n\ndef is_in_val_zone(location, val_location, val_zone_size, BUFFER=RandomOpt().BUFFER):\n    x = location[0]\n    y = location[1]\n    x_match = val_location[0] - BUFFER <= x <= val_location[0] + val_zone_size[0] + BUFFER\n    y_match = val_location[1] - BUFFER <= y <= val_location[1] + val_zone_size[1] + BUFFER\n    return x_match and y_match\n\nclass RandomPatchLocDataset(data.Dataset):\n    def __init__(self, mask, val_location, val_zone_size):\n        self.mask = mask\n        self.val_location = val_location\n        self.val_zone_size = val_zone_size\n        self.sample_random_location_train = lambda x: sample_random_location(mask.shape)\n        self.is_in_mask_train = lambda x: is_in_masked_zone(x, mask)\n\n    def is_proper_train_location(self, location):\n        return not is_in_val_zone(location, self.val_location, self.val_zone_size) and self.is_in_mask_train(location)\n\n    def __len__(self):\n        return 1280\n\n    def __getitem__(self, index):\n        # Generate a random patch\n        # Ignore the index\n        loc = self.sample_random_location_train(0)\n        while not self.is_proper_train_location(loc):\n            loc = self.sample_random_location_train(0)\n        return loc.int().squeeze(1)","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.736308,"end_time":"2023-03-31T17:08:31.240545","exception":false,"start_time":"2023-03-31T17:08:30.504237","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T07:54:26.116581Z","iopub.execute_input":"2023-04-02T07:54:26.117007Z","iopub.status.idle":"2023-04-02T07:54:26.993960Z","shell.execute_reply.started":"2023-04-02T07:54:26.116969Z","shell.execute_reply":"2023-04-02T07:54:26.992825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define training and validation model\nThe full code with the model with Unet and random patch.\n\nWe will save the best model on validation.\n\nNOTE: You can also download the model in OUTPUT and run test Code.","metadata":{"papermill":{"duration":0.006092,"end_time":"2023-03-31T17:08:31.252466","exception":false,"start_time":"2023-03-31T17:08:31.246374","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ============= Model ==============\nclass ModelOpt:\n    def __init__(self):\n        # self.GPU_ID = '0'  \n        self.Z_DIM = RandomOpt().Z_DIM\n        self.BUFFER = RandomOpt().BUFFER\n        self.SEED = 0\n        self.BATCH_SIZE = 64\n        self.LEARNING_RATE =1e-4\n        self.TRAINING_EPOCH = 25\n        self.LOG_DIR = '../working'\n        self.LOAD_VOLUME = [1, 2, 3]\n        # Val\n        self.VAL_LOC = (1300, 1000)\n        self.VAL_SIZE = (300, 7000)\n\nclass RandomPatchModel():\n    def __init__(self, opt = ModelOpt()):\n        self.opt = opt\n        self._setup_all()\n        self.volume_list = [load_volume('train', i) for i in opt.LOAD_VOLUME]\n        # Here volume: [Z_DIM, SHARED_HEIGHT, W_V1 + W_V2 + ...]\n        self.volume = torch.cat(self.volume_list, dim=2)\n        # Same for mask and label\n        self.mask_list = [load_mask('train', i) for i in opt.LOAD_VOLUME]\n        self.labels_list = [load_labels('train', i) for i in opt.LOAD_VOLUME]\n        # [SHARED_HEIGHT, W_V1 + W_V2 + ...]\n        self.labels = torch.cat(self.labels_list, dim=1)\n        self.mask = torch.cat(self.mask_list, dim=1)\n\n        self.net = create_model(opt.Z_DIM).to(self.device)\n\n        # Dataset\n        self.loc_datast = RandomPatchLocDataset(self.mask, val_location=opt.VAL_LOC, val_zone_size=opt.VAL_SIZE)\n        self.loc_loader = data.DataLoader(self.loc_datast, batch_size=opt.BATCH_SIZE)\n        # Val\n        self.val_loc = []\n        for x in range(opt.VAL_LOC[0], opt.VAL_LOC[0] + opt.VAL_SIZE[0], opt.BUFFER):\n            for y in range(opt.VAL_LOC[1], opt.VAL_LOC[1] + opt.VAL_SIZE[1], opt.BUFFER):\n                if is_in_masked_zone([torch.tensor(x),torch.tensor(y)], self.mask):\n                    self.val_loc.append([[x, y]])\n        print(f\"======> Num Patches Val: {len(self.val_loc)}\")\n\n\n    def _setup_all(self):\n        # random seed\n        np.random.seed(self.opt.SEED)\n        torch.manual_seed(self.opt.SEED)\n        torch.cuda.manual_seed_all(self.opt.SEED)\n        # torch\n        # os.environ['CUDA_VISIBLE_DEVICES'] = self.opt.GPU_ID\n        torch.backends.cudnn.enabled = True\n        torch.backends.cudnn.benchmark = True\n        self.device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n        # Log\n        self.log_dir = self.opt.LOG_DIR\n        self.ckpt = os.path.join(self.log_dir)\n\n    def get_subvolume(self, batch_loc, volume, labels):\n        # batch_loc : [batch_size, 2]\n        subvolume = []\n        label = []\n        for l in batch_loc:\n            x = l[0]\n            y = l[1]\n            sv = volume[:, x - self.opt.BUFFER:x + self.opt.BUFFER, y - self.opt.BUFFER:y + self.opt.BUFFER]\n            sv = sv / 65535.\n            subvolume.append(sv)\n            if labels is not None:\n                lb = labels[x - self.opt.BUFFER:x + self.opt.BUFFER, y - self.opt.BUFFER:y + self.opt.BUFFER]\n                lb = lb.unsqueeze(0)\n                label.append(lb)\n        # [batch, Z_DIM, BUFFER, BUFFER]\n        subvolume = torch.stack(subvolume)\n        # [batch, 1, BUFFER, BUFFER]\n        if labels is not None:\n            label = torch.stack(label)\n        return subvolume, label\n\n    def augment_train_data(self, subvolume, label):\n        # Add Data augmentation here\n        return subvolume, label\n\n    def train_loop(self):\n        print(\"=====> Begin training\")\n#         self.criterion = torch.nn.BCEWithLogitsLoss(reduction='mean')\n        self.criterion = calc_loss\n        self.optimizer = optim.Adam(self.net.parameters(), lr=self.opt.LEARNING_RATE)\n        self.net.train()\n\n        best_val_loss = 100\n        best_val_acc = 0\n        meter = AverageMeter()\n        for epoch in range(self.opt.TRAINING_EPOCH):\n            bar = tqdm(enumerate(self.loc_loader), total=len(self.loc_datast) / self.opt.BATCH_SIZE)\n            bar.set_description_str(f\"Epoch: {epoch}\")\n            for i, loc in bar:\n                subvolume, label = self.get_subvolume(loc, self.volume, self.labels)\n                loss = self._train_step(subvolume, label)\n                meter.update(loss)\n                bar.set_postfix_str(f\"Avg loss: {np.round(meter.get_value(),3)}\")\n\n            val_loss, val_acc = self.validataion_loop()\n            print(f\"======> Val Loss:{np.round(val_loss,3)} | Val Acc:{np.round(val_acc,3)} \")\n            if val_loss < best_val_loss and val_acc > best_val_acc:\n                torch.save(self.net.state_dict(), os.path.join(self.ckpt, \"best.pt\"))\n                print(\"======> Save best val model\")\n\n                best_val_loss = val_loss\n                best_val_acc = val_acc\n\n\n\n    def _train_step(self, subvolume, label):\n        self.optimizer.zero_grad()\n        # inputs: subvolume: [batch, Z_DIM, BUFFER, BUFFER]\n        #         label: [batch, 1, BUFFER, BUFFER]\n        outputs = self.net(subvolume.to(self.device))\n        loss = self.criterion(outputs, label.to(self.device))\n        loss.backward()\n        self.optimizer.step()\n        return loss\n\n    def validataion_loop(self):\n        meter_loss = AverageMeter()\n        meter_acc = AverageMeter()\n        self.net.eval()\n        for loc in self.val_loc:\n            subvolume, label = self.get_subvolume(loc, self.volume, self.labels)\n            outputs = self.net(subvolume.to(self.device))\n            loss = self.criterion(outputs, label.to(self.device))\n            meter_loss.update(loss)\n            pred = torch.sigmoid(outputs) > THRESHOLD #0.5\n            meter_acc.update(\n                (pred == label.to(self.device)).sum(),\n                int(torch.prod(torch.tensor(label.shape)))\n            )\n        self.net.train()\n        return meter_loss.get_value(), meter_acc.get_value()\n\n    def load_best_ckpt(self):\n        self.net.load_state_dict(torch.load(os.path.join(self.ckpt, \"best.pt\")))\n#        self.net.load_state_dict(torch.load(BEST_MODEL_PATH, map_location='cuda'))\n\n\n# For the metric\nclass AverageMeter(object):\n    def __init__(self):\n        self.sum = 0\n        self.n = 0\n\n    def update(self, x, n=1):\n        self.sum += float(x)\n        self.n += n\n\n    def reset(self):\n        self.sum = 0\n        self.n = 0\n\n    def get_value(self):\n        if self.n:\n            return self.sum / self.n\n        return 0","metadata":{"papermill":{"duration":0.037184,"end_time":"2023-03-31T17:08:31.296051","exception":false,"start_time":"2023-03-31T17:08:31.258867","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T07:54:26.995887Z","iopub.execute_input":"2023-04-02T07:54:26.996300Z","iopub.status.idle":"2023-04-02T07:54:27.027761Z","shell.execute_reply.started":"2023-04-02T07:54:26.996258Z","shell.execute_reply":"2023-04-02T07:54:27.026584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Define model\nmodel = RandomPatchModel()","metadata":{"papermill":{"duration":219.055368,"end_time":"2023-03-31T17:12:10.356806","exception":false,"start_time":"2023-03-31T17:08:31.301438","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T07:54:27.029427Z","iopub.execute_input":"2023-04-02T07:54:27.029896Z","iopub.status.idle":"2023-04-02T07:57:36.538070Z","shell.execute_reply.started":"2023-04-02T07:54:27.029856Z","shell.execute_reply":"2023-04-02T07:57:36.536828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training\nif is_train:\n    model.train_loop()","metadata":{"execution":{"iopub.status.busy":"2023-04-02T07:57:36.540807Z","iopub.execute_input":"2023-04-02T07:57:36.541742Z","iopub.status.idle":"2023-04-02T08:03:52.354114Z","shell.execute_reply.started":"2023-04-02T07:57:36.541708Z","shell.execute_reply":"2023-04-02T08:03:52.352615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Validation","metadata":{"papermill":{"duration":0.099932,"end_time":"2023-03-31T17:21:14.298760","exception":false,"start_time":"2023-03-31T17:21:14.198828","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import torch\n\nif is_train:\n    # Load the best model\n    model.load_best_ckpt()\n    # model.criterion = torch.nn.BCEWithLogitsLoss(reduction='mean')\n    model.criterion  = calc_loss\n    loss, acc = model.validataion_loop()\n    model.net.eval()\n    print(f\"Val loss: {np.round(loss,3)} | Val acc: {np.round(acc, 3)}\")","metadata":{"papermill":{"duration":7.937365,"end_time":"2023-03-31T17:21:22.335760","exception":false,"start_time":"2023-03-31T17:21:14.398395","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T08:03:52.357122Z","iopub.execute_input":"2023-04-02T08:03:52.357428Z","iopub.status.idle":"2023-04-02T08:03:59.281955Z","shell.execute_reply.started":"2023-04-02T08:03:52.357399Z","shell.execute_reply":"2023-04-02T08:03:59.280784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"# is infer:\n#     # Load the best model\n#     model = model.load_best_ckpt()\n#     # model.criterion = torch.nn.BCEWithLogitsLoss(reduction='mean')\n#     model.criterion  = calc_loss\n#     # set evaluation_mode\n#     model = model.net.eval()","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-04-02T08:03:59.286069Z","iopub.execute_input":"2023-04-02T08:03:59.287205Z","iopub.status.idle":"2023-04-02T08:03:59.291999Z","shell.execute_reply.started":"2023-04-02T08:03:59.287156Z","shell.execute_reply":"2023-04-02T08:03:59.290743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\ndef compute_predictions_map(split, index):\n    print(f\"======> Load data for {split}/{index}\")\n    test_volume = load_volume(split=split, index=index)\n    test_mask = load_mask(split=split, index=index)\n    print(f\"======> Volume shape: {test_volume.shape}\")\n    test_locations = []\n    BUFFER = model.opt.BUFFER\n    stride = BUFFER // 2\n\n    for x in range(BUFFER, test_volume.shape[1] - BUFFER, stride):\n        for y in range(BUFFER, test_volume.shape[2] - BUFFER, stride):\n            if is_in_masked_zone([torch.tensor(x),torch.tensor(y)], test_mask):\n                test_locations.append((x, y))\n    print(f\"======> {len(test_locations)} test locations (after filtering by mask)\")\n\n    predictions_map = torch.zeros((1, 1, test_volume.shape[1], test_volume.shape[2]))\n    predictions_map_counts = torch.zeros((1, 1, test_volume.shape[1], test_volume.shape[2]))\n    print(f\"======> Compute predictions\")\n\n    with torch.no_grad():\n        bar = tqdm(test_locations)\n        for loc in bar:\n            subvolume, label = model.get_subvolume([loc], test_volume, None)\n            outputs = model.net(subvolume.to(model.device))\n            pred = torch.sigmoid(outputs)\n            # print(loc, (pred > 0.5).sum())\n            # Here a single location may be with multiple result\n            predictions_map[:, :, loc[0] - BUFFER : loc[0] + BUFFER, loc[1] - BUFFER : loc[1] + BUFFER] += pred.cpu()\n            predictions_map_counts[:, :, loc[0] - BUFFER : loc[0] + BUFFER, loc[1] - BUFFER : loc[1] + BUFFER] += 1\n\n    # print(predictions_map_b[:,:, 2500, 1000])\n    # print(predictions_map_counts[:,:, 2500, 1000])\n    predictions_map /= (predictions_map_counts + 1e-7)\n    return predictions_map","metadata":{"papermill":{"duration":0.093289,"end_time":"2023-03-31T17:21:22.505344","exception":false,"start_time":"2023-03-31T17:21:22.412055","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T08:03:59.293535Z","iopub.execute_input":"2023-04-02T08:03:59.295796Z","iopub.status.idle":"2023-04-02T08:03:59.308178Z","shell.execute_reply.started":"2023-04-02T08:03:59.295745Z","shell.execute_reply":"2023-04-02T08:03:59.307040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n# Threshold is very important !!!!!\n# plt.imshow(predictions_map_a.squeeze() > 0.15, cmap='gray')","metadata":{"papermill":{"duration":2.461975,"end_time":"2023-03-31T17:32:51.858748","exception":false,"start_time":"2023-03-31T17:32:49.396773","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T08:03:59.309859Z","iopub.execute_input":"2023-04-02T08:03:59.310231Z","iopub.status.idle":"2023-04-02T08:03:59.321150Z","shell.execute_reply.started":"2023-04-02T08:03:59.310192Z","shell.execute_reply":"2023-04-02T08:03:59.320048Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plt.imshow(predictions_map_b.squeeze() > 0.15, cmap='gray')","metadata":{"papermill":{"duration":2.201709,"end_time":"2023-03-31T17:32:54.422198","exception":false,"start_time":"2023-03-31T17:32:52.220489","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T08:03:59.322738Z","iopub.execute_input":"2023-04-02T08:03:59.323133Z","iopub.status.idle":"2023-04-02T08:03:59.330756Z","shell.execute_reply.started":"2023-04-02T08:03:59.323095Z","shell.execute_reply":"2023-04-02T08:03:59.329431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission\n\nRescale the pred to the raw size and create 'submission.csv'","metadata":{"papermill":{"duration":0.43222,"end_time":"2023-03-31T17:32:55.245110","exception":false,"start_time":"2023-03-31T17:32:54.812890","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from skimage.transform import resize as resize_ski\nimport PIL.Image as Image\n\ndef rle(predictions_map, threshold):\n    flat_img = predictions_map.flatten()\n    flat_img = np.where(flat_img > threshold, 1, 0).astype(np.uint8)\n\n    starts = np.array((flat_img[:-1] == 0) & (flat_img[1:] == 1))\n    ends = np.array((flat_img[:-1] == 1) & (flat_img[1:] == 0))\n    starts_ix = np.where(starts)[0] + 2\n    ends_ix = np.where(ends)[0] + 2\n    lengths = ends_ix - starts_ix\n    return \" \".join(map(str, sum(zip(starts_ix, lengths), ())))","metadata":{"execution":{"iopub.status.busy":"2023-04-02T08:03:59.332467Z","iopub.execute_input":"2023-04-02T08:03:59.332845Z","iopub.status.idle":"2023-04-02T08:03:59.860906Z","shell.execute_reply.started":"2023-04-02T08:03:59.332807Z","shell.execute_reply":"2023-04-02T08:03:59.859815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sorted(os.listdir(\"../input/vesuvius-challenge-ink-detection/test\"))","metadata":{"execution":{"iopub.status.busy":"2023-04-02T08:20:01.513391Z","iopub.execute_input":"2023-04-02T08:20:01.514512Z","iopub.status.idle":"2023-04-02T08:20:01.524060Z","shell.execute_reply.started":"2023-04-02T08:20:01.514467Z","shell.execute_reply":"2023-04-02T08:20:01.522568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if is_infer:\n    # load test data\n    # load test folders\n    DATA_DIR = \"../input/vesuvius-challenge-ink-detection\"\n    TEST_DIRS = sorted(os.listdir(\"../input/vesuvius-challenge-ink-detection/test\"))\n\n    predictions_maps = []\n    for index in TEST_DIRS:\n        predictions_map = compute_predictions_map(split=\"test\", index=index)\n        predictions_maps.append(predictions_map)\n        \n    original_size_imgs = []\n    for TEST_DIR in TEST_DIRS:\n        original_size_img = Image.open(DATA_DIR + f\"/test/{TEST_DIR}/mask.png\").size\n        original_size_imgs.append(original_size_img)   \n\n    rescaled_predictions_maps = []\n    for prediction_map, original_size_img in zip(predictions_maps, original_size_imgs):\n        prediction_map = resize_ski(prediction_map.squeeze(), original_size_img).squeeze()\n        rescaled_predictions_maps.append(prediction_map)\n   \n    rles = []\n    for rescaled_predictions_map in rescaled_predictions_maps:\n        rle_value = rle(rescaled_predictions_map, threshold=THRESHOLD)\n        rles.append(rle_value)\n\n    import pandas as pd\n    submission = pd.DataFrame({'Id': TEST_DIRS,\n                               'Predicted': rles})\n\n    submission.to_csv('../working/submission.csv', index=False)","metadata":{"tags":[],"execution":{"iopub.status.busy":"2023-04-02T08:20:16.639752Z","iopub.execute_input":"2023-04-02T08:20:16.640412Z","iopub.status.idle":"2023-04-02T08:20:30.258710Z","shell.execute_reply.started":"2023-04-02T08:20:16.640372Z","shell.execute_reply":"2023-04-02T08:20:30.257210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Next step\nYou can\n- Try different architecture in Unet:\n    - Other size\n    - Attention Block\n- Tune the hyperparameter\n    - Z_DIM ...\n- Use more suitable segmentation Loss","metadata":{"papermill":{"duration":0.442124,"end_time":"2023-03-31T17:48:35.503632","exception":false,"start_time":"2023-03-31T17:48:35.061508","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# END","metadata":{"papermill":{"duration":0.368616,"end_time":"2023-03-31T17:48:36.249776","exception":false,"start_time":"2023-03-31T17:48:35.881160","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-04-02T08:14:36.989896Z","iopub.execute_input":"2023-04-02T08:14:36.990625Z","iopub.status.idle":"2023-04-02T08:14:36.996507Z","shell.execute_reply.started":"2023-04-02T08:14:36.990577Z","shell.execute_reply":"2023-04-02T08:14:36.995433Z"},"trusted":true},"execution_count":null,"outputs":[]}]}