{"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":"IGNORE_FRAGMENT = \"none\"\nQUICK_SAVE = True\n\nDEBUG = False\nIS_ROT_TTA = True\nIS_FLIP_TTA = False\nTHR = 0.625\n\nMODEL_PATH = \"/kaggle/input/\"","metadata":{"execution":{"iopub.status.busy":"2023-06-13T22:40:49.149599Z","iopub.execute_input":"2023-06-13T22:40:49.150054Z","iopub.status.idle":"2023-06-13T22:40:49.155798Z","shell.execute_reply.started":"2023-06-13T22:40:49.150023Z","shell.execute_reply":"2023-06-13T22:40:49.154435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install --quiet /kaggle/input/segmentation-models-pytorch-offline-installer/safetensors-0.3.1-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install --quiet /kaggle/input/segmentation-models-pytorch-offline-installer/timm-0.9.2-py3-none-any.whl\n!pip install --quiet /kaggle/input/segmentation-models-pytorch-offline-installer/pretrainedmodels-0.7.4.tar.gz\n!pip install --quiet /kaggle/input/segmentation-models-pytorch-offline-installer/efficientnet_pytorch-0.7.1.tar.gz\n!pip install --quiet /kaggle/input/segmentation-models-pytorch-offline-installer/segmentation_models_pytorch-0.3.3-py3-none-any.whl","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-13T22:40:49.157890Z","iopub.execute_input":"2023-06-13T22:40:49.158560Z","iopub.status.idle":"2023-06-13T22:43:27.702359Z","shell.execute_reply.started":"2023-06-13T22:40:49.158526Z","shell.execute_reply":"2023-06-13T22:43:27.701141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport filecmp\nimport gc\nimport numpy as np\nimport os\nimport pandas as pd\nimport random\nimport segmentation_models_pytorch as smp\nimport PIL.Image as Image\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom collections import defaultdict\nfrom pathlib import Path\nfrom tqdm.notebook import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-13T22:43:27.704449Z","iopub.execute_input":"2023-06-13T22:43:27.705153Z","iopub.status.idle":"2023-06-13T22:43:27.712606Z","shell.execute_reply.started":"2023-06-13T22:43:27.705115Z","shell.execute_reply":"2023-06-13T22:43:27.711619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append(\"/kaggle/input/unet3d-partial-conv\")\nfrom model_list import UNet3D, ResidualUNetSE3D, SegFormer","metadata":{"execution":{"iopub.status.busy":"2023-06-13T22:43:27.714980Z","iopub.execute_input":"2023-06-13T22:43:27.715271Z","iopub.status.idle":"2023-06-13T22:43:27.724601Z","shell.execute_reply.started":"2023-06-13T22:43:27.715248Z","shell.execute_reply":"2023-06-13T22:43:27.723653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\ndef load_images(input_dir, img_depth, z_start, z_end, resize_ratio):\n\n    # all images\n    image_path_list = sorted(list(Path(input_dir / \"surface_volume\").glob('*.tif')))\n    # select candidates\n    image_path_list = image_path_list[z_start:z_end + img_depth]\n\n    img_stack = []\n    for image_path in tqdm(image_path_list, total=len(image_path_list), dynamic_ncols=True, desc=\"Loading images\"):\n        img = np.array(Image.open(str(image_path)), dtype=np.uint16)\n        img = (img >> 8).astype(np.uint8)\n        img = cv2.resize(img, dsize=None, fx=1.0/resize_ratio, fy=1.0/resize_ratio, interpolation=cv2.INTER_AREA)\n        img_stack.append(img)\n        del img\n        gc.collect()\n\n    img_stack = np.stack(img_stack, axis=0)\n\n    return img_stack\n\ndef get_scanning_positions(mask, roi_size, stride, offset=0):\n\n    h, w = mask.shape\n    x_pos_list = []\n    y_pos_list = []\n    for y in range(offset, h-roi_size+1, stride):\n        for x in range(offset, w-roi_size+1, stride):\n            if mask[y:y+roi_size, x:x+roi_size].mean() > 0.1:\n                x_pos_list.append(x)\n                y_pos_list.append(y)\n\n    return x_pos_list, y_pos_list\n\ndef process_tile(model, device, img_tile, is_rot_tta, is_flip_tta):\n\n    img_tile = torch.as_tensor(img_tile, dtype=torch.float32).unsqueeze(0).to(device)\n\n    predictions = []\n    with torch.no_grad():\n        pred = model(img_tile)\n        predictions.append(pred.cpu())\n\n        if is_rot_tta:\n            # rotation TTA\n            for i in range(1,4):\n                pred = model(torch.rot90(img_tile, k=i, dims=(-2, -1)))\n                pred = torch.rot90(pred, k=-i, dims=(-2, -1)) # Rotate back\n                predictions.append(pred.cpu())\n\n        if is_flip_tta:\n            # flip TTA\n            pred = model(torch.flip(img_tile, dims=(-2,)))\n            pred = torch.flip(pred, dims=(-2,)) # Flip back\n            predictions.append(pred.cpu())\n\n            pred = model(torch.flip(img_tile, dims=(-1,)))\n            pred = torch.flip(pred, dims=(-1,)) # Flip back\n            predictions.append(pred.cpu())\n\n    predictions = torch.stack(predictions).numpy()\n\n    return np.mean(predictions, axis=0)\n\ndef normalize_img(img):\n\n    img = img.astype(np.float32)\n    if img.sum() != 0:\n        img[img > 0] = (img[img > 0] - img[img > 0].mean()) / (img[img > 0].std() + 1e-8)\n\n    return img\n\ndef inference(\n    img_stack, mask_resized,\n    z_start, z_end, img_depth,\n    num_instances,\n    roi_size, stride,\n    model, device,\n    is3d,\n    is_rot_tta, is_flip_tta,\n):\n    # Get the scanning position of the image\n    x_pos_list, y_pos_list = get_scanning_positions(mask_resized, roi_size, stride)\n\n    # inference\n    _, input_h, input_w = img_stack.shape\n\n    output_img = np.zeros((input_h, input_w), dtype=np.float32)\n    count_img = np.zeros((input_h, input_w), dtype=np.float32)\n\n    weight = np.ones((roi_size, roi_size), dtype=np.float32)\n    weight[0:roi_size//4, :] *= np.linspace(0, 1, num=roi_size//4, dtype=np.float32).reshape(-1, 1)\n    weight[-roi_size//4:, :] *= np.linspace(1, 0, num=roi_size//4, dtype=np.float32).reshape(-1, 1)\n    weight[:, 0:roi_size//4] *= np.linspace(0, 1, num=roi_size//4, dtype=np.float32).reshape(1, -1)\n    weight[:, -roi_size//4:] *= np.linspace(1, 0, num=roi_size//4, dtype=np.float32).reshape(1, -1)\n\n    for y, x in tqdm(zip(y_pos_list, x_pos_list), total=len(y_pos_list), dynamic_ncols=True, desc=\"Inference\"):\n\n        z_start_list = np.linspace(0, z_end-z_start, num=num_instances, dtype=np.int32)\n\n        img_tile = []\n        for z in z_start_list:\n            img = img_stack[z:z + img_depth, y:y + roi_size, x:x + roi_size]\n            img = (img / 255.0).astype(np.float32)\n            if is3d:\n                img = img[np.newaxis, ...] # DHW -> CDHW\n            img_tile.append(normalize_img(img))\n        img_tile = np.stack(img_tile, axis=0)\n\n        output = process_tile(model, device, img_tile, is_rot_tta, is_flip_tta)\n\n        output_img[y:y+roi_size, x:x+roi_size] += output * weight\n        count_img[y:y+roi_size, x:x+roi_size] += weight\n\n    # averaging\n    output_img = output_img / np.maximum(count_img, 1e-5)\n\n    torch.cuda.empty_cache()\n\n    return output_img","metadata":{"execution":{"iopub.status.busy":"2023-06-13T22:43:27.728066Z","iopub.execute_input":"2023-06-13T22:43:27.728388Z","iopub.status.idle":"2023-06-13T22:43:27.757114Z","shell.execute_reply.started":"2023-06-13T22:43:27.728364Z","shell.execute_reply":"2023-06-13T22:43:27.756134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference_fold(\n    fold_id,\n    fragment_path,\n    model_dir,\n    device,\n    is_rot_tta, is_flip_tta,\n):\n\n    mask_org = cv2.imread(str(fragment_path / \"mask.png\"), cv2.IMREAD_GRAYSCALE)\n    mask_org[mask_org>0] = 1\n\n    #-------------------------------\n    # stage1: CNN based segmentation\n    print(f\"fold{fold_id}: stage1\")\n    resize_ratio = 2\n    z_start = 20\n    z_end = 55\n    img_depth = 8\n    num_instances = 12\n\n    roi_size = 384\n    stride = roi_size // 2\n\n    # weight is calculated by forward selection\n    weight0 = 0.49\n\n    # loading\n    img_stack = load_images(fragment_path, img_depth, z_start, z_end, resize_ratio)\n    mask_resized = cv2.resize(mask_org, dsize=None, fx=1.0/resize_ratio, fy=1.0/resize_ratio, interpolation=cv2.INTER_NEAREST)\n\n    # padding\n    img_stack = np.pad(img_stack, [(0, 0), (0, stride), (0, stride)], 'constant')\n    mask_resized = np.pad(mask_resized, [(0, stride), (0, stride)], 'constant')\n\n    output_map = np.zeros_like(mask_resized, dtype=np.float32)\n\n    ###\n    # UNet3D\n    print(f\"UNet3D\")\n    model = UNet3D()\n    weight = torch.load(f\"{model_dir}/vcid-exp545/fold{fold_id}/val_dice.ckpt\")\n    model.load_state_dict(weight[\"state_dict\"], strict=False)\n    model = model.to(device).eval()\n\n    pred = inference(\n        img_stack, mask_resized,\n        z_start, z_end, img_depth,\n        num_instances, roi_size, stride,\n        model, device,\n        is3d=True,\n        is_rot_tta=is_rot_tta,\n        is_flip_tta=is_flip_tta,\n    )\n    output_map = pred\n\n    del model, pred\n    gc.collect()\n    torch.cuda.empty_cache()\n    # UNet3D\n    ###\n\n    ###\n    # ResidualUNetSE3D\n    print(f\"ResidualUNetSE3D\")\n    model = ResidualUNetSE3D()\n    weight = torch.load(f\"{model_dir}/vcid-exp544/fold{fold_id}/val_dice.ckpt\")\n    model.load_state_dict(weight[\"state_dict\"], strict=False)\n    model = model.to(device).eval()\n\n    pred = inference(\n        img_stack, mask_resized,\n        z_start, z_end, img_depth,\n        num_instances, roi_size, stride,\n        model, device,\n        is3d=True,\n        is_rot_tta=is_rot_tta,\n        is_flip_tta=is_flip_tta,\n    )\n    output_map = (1 - weight0) * output_map + weight0 * pred\n\n    del model, pred\n    gc.collect()\n    torch.cuda.empty_cache()\n    # ResidualUNetSE3D\n    ###\n\n    # resize\n    output_map = cv2.resize(output_map, dsize=None, fx=resize_ratio, fy=resize_ratio, interpolation=cv2.INTER_CUBIC)\n    output_map = np.clip(output_map, 0, 1)\n    # masking\n    output_map = output_map[:mask_org.shape[0], :mask_org.shape[1]] * (mask_org > 0)\n\n    del img_stack, mask_resized\n    gc.collect()\n\n    # stage1: CNN based segmentation\n    #-------------------------------\n\n    #-------------------------------\n    # stage2: Transformer based segmentation\n    print(f\"fold{fold_id}: stage2\")\n    resize_ratio = 1\n    z_start = 10\n    z_end = 55\n    img_depth = 3\n    num_instances = 20\n\n    roi_size = 384\n    stride = roi_size // 2\n\n    stage1_thr = 0.57\n\n    # loading\n    img_stack = load_images(fragment_path, img_depth, z_start, z_end, resize_ratio)\n    mask = np.zeros_like(output_map)\n    mask[output_map>stage1_thr] = 1\n    mask_resized = cv2.resize(mask, dsize=None, fx=1.0/resize_ratio, fy=1.0/resize_ratio, interpolation=cv2.INTER_NEAREST)\n\n    # padding\n    img_stack = np.pad(img_stack, [(0, 0), (0, stride), (0, stride)], 'constant')\n    mask_resized = np.pad(mask_resized, [(0, stride), (0, stride)], 'constant')\n\n    print(f\"SegFormer\")\n    model = SegFormer()\n    weight = torch.load(f\"{model_dir}/vcid-exp568/fold{fold_id}/val_dice.ckpt\")\n    model.load_state_dict(weight[\"state_dict\"], strict=False)\n    model = model.to(device).eval()\n\n    output_map = inference(\n        img_stack, mask_resized,\n        z_start, z_end, img_depth,\n        num_instances, roi_size, stride,\n        model, device,\n        is3d=False,\n        is_rot_tta=is_rot_tta,\n        is_flip_tta=is_flip_tta,\n    )\n\n    del model\n    gc.collect()\n    torch.cuda.empty_cache()\n\n    # resize\n    output_map = cv2.resize(output_map, dsize=None, fx=resize_ratio, fy=resize_ratio, interpolation=cv2.INTER_CUBIC)\n    output_map = np.clip(output_map, 0, 1)\n    # masking\n    output_map = output_map[:mask_org.shape[0], :mask_org.shape[1]] * (mask_org > 0)\n\n    del img_stack, mask, mask_resized\n    gc.collect()\n\n    # stage2: Transformer based segmentation\n    #-------------------------------\n\n    del mask_org\n    gc.collect()\n\n    return output_map","metadata":{"execution":{"iopub.status.busy":"2023-06-13T22:43:27.758349Z","iopub.execute_input":"2023-06-13T22:43:27.758728Z","iopub.status.idle":"2023-06-13T22:43:27.781587Z","shell.execute_reply.started":"2023-06-13T22:43:27.758697Z","shell.execute_reply":"2023-06-13T22:43:27.780569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle(img):\n    flat_img = img.flatten().astype(np.uint8)\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-06-13T22:43:27.782751Z","iopub.execute_input":"2023-06-13T22:43:27.783129Z","iopub.status.idle":"2023-06-13T22:43:27.795474Z","shell.execute_reply.started":"2023-06-13T22:43:27.783098Z","shell.execute_reply":"2023-06-13T22:43:27.794572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(42)\n\nsample_submission_flag = filecmp.cmp(\n    \"../input/vesuvius-challenge-ink-detection/test/a/surface_volume/00.tif\",\n    \"../input/vcid-file-check/00.tif\",\n    shallow=True\n)\n\nif sample_submission_flag and QUICK_SAVE:\n    df_sub = pd.read_csv(\"../input/vesuvius-challenge-ink-detection/sample_submission.csv\")\n    df_sub.to_csv(\"submission.csv\", index=False)\nelse:\n\n    device = torch.device(\"cuda:0\")\n\n    input_dir = Path(\"../input/vesuvius-challenge-ink-detection/test\")\n    test_fragments = sorted([ fragment_name for fragment_name in input_dir.iterdir()])\n\n    pred_images = []\n    for test_fragment_path in test_fragments:\n\n        mask = cv2.imread(str(test_fragment_path / \"mask.png\"), cv2.IMREAD_GRAYSCALE)\n        if IGNORE_FRAGMENT in test_fragment_path.name:\n            pred_image = np.ones_like(mask, dtype=np.uint8) * (mask > 0)\n            pred_images.append(rle(pred_image))\n            del mask, pred_image\n            gc.collect()\n            continue\n\n        predictions = []\n        for fold_id in range(0, 5):\n\n            pred = inference_fold(\n                fold_id=fold_id,\n                fragment_path=test_fragment_path,\n                model_dir=MODEL_PATH,\n                device=device,\n                is_rot_tta=IS_ROT_TTA,\n                is_flip_tta=IS_FLIP_TTA,\n            )\n            predictions.append(pred)\n\n        predictions = np.array(predictions)\n        output = np.mean(predictions, axis=0)\n\n        pred_image = np.zeros(output.shape, dtype=np.uint8)\n        pred_image[output > THR] = 1\n\n        pred_images.append(rle(pred_image))\n\n        if DEBUG:\n            job_type = __file__.split(\"/\")[-1].split(\".\")[0].split(\"_\")[0]\n            cv2.imwrite(f\"{job_type}_{test_fragment_path.name}.png\", (output * 255).astype(np.uint8))\n            cv2.imwrite(f\"{job_type}_{test_fragment_path.name}_thr.png\", (pred_image * 255).astype(np.uint8))\n\n    submission = defaultdict(list)\n    for fragment_id, fragment_name in enumerate(test_fragments):\n        submission[\"Id\"].append(fragment_name.name)\n        submission[\"Predicted\"].append(pred_images[fragment_id])\n\n    pd.DataFrame.from_dict(submission).to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-06-13T22:43:27.796777Z","iopub.execute_input":"2023-06-13T22:43:27.797470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub = pd.read_csv(\"submission.csv\")\nprint(df_sub.head())","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}