{"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,glob\nimport filecmp\nimport gc\nimport cv2\nimport sys\nimport random\nimport imageio\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom IPython.display import Video\nimport tifffile as tiff\nimport torch\nimport torchvision\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda import amp\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\nsys.path.append(\"/kaggle/input/resnet3d\")\nfrom resnet3d import generate_model","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"QUICK_SAVE = False \nmode = 'test'\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nchkps = [\n    '/kaggle/input/train-r18-3d-mosaic/vesuvius-models/Resnet18_3d_fold1_best.pth',\n    '/kaggle/input/train-r18-3d-mosaic2-exp1/vesuvius-models/Resnet18_3d_fold2a_best.pth',\n    '/kaggle/input/train-r34-3d-mosaic/vesuvius-models/Resnet34_3d_fold2b_best.pth',\n    '/kaggle/input/train-r34-3d-mosaic/vesuvius-models/Resnet34_3d_fold1_best.pth',\n    '/kaggle/input/train-r34-3d-mosaic/vesuvius-models/Resnet34_3d_fold3_best.pth',\n    '/kaggle/input/train-r34-3d-mosaic/vesuvius-models/Resnet34_3d_fold2a_best.pth']\ncomp_dataset_path='/kaggle/input/vesuvius-challenge-ink-detection/'\nTH=0.5\nexpand=128\nsize=896\ntile_size=size+expand\nprint('tile_size=',tile_size)\nstride=tile_size//2     \nBATCH_SIZE = 1\nZ_START = 16\nZ_DIMS = 32","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cupy as cp\nxp = cp\n\ndelta_lookup = {\n    \"xx\": xp.array([[1, -2, 1]], dtype=float),\n    \"yy\": xp.array([[1], [-2], [1]], dtype=float),\n    \"xy\": xp.array([[1, -1], [-1, 1]], dtype=float),\n}\n\ndef operate_derivative(img_shape, pair):\n    assert len(img_shape) == 2\n    delta = delta_lookup[pair]\n    fft = xp.fft.fftn(delta, img_shape)\n    return fft * xp.conj(fft)\n\ndef soft_threshold(vector, threshold):\n    return xp.sign(vector) * xp.maximum(xp.abs(vector) - threshold, 0)\n\ndef back_diff(input_image, dim):\n    assert dim in (0, 1)\n    r, n = xp.shape(input_image)\n    size = xp.array((r, n))\n    position = xp.zeros(2, dtype=int)\n    temp1 = xp.zeros((r+1, n+1), dtype=float)\n    temp2 = xp.zeros((r+1, n+1), dtype=float)\n    \n    temp1[position[0]:size[0], position[1]:size[1]] = input_image\n    temp2[position[0]:size[0], position[1]:size[1]] = input_image\n    \n    size[dim] += 1\n    position[dim] += 1\n    temp2[position[0]:size[0], position[1]:size[1]] = input_image\n    temp1 -= temp2\n    size[dim] -= 1\n    return temp1[0:size[0], 0:size[1]]\n\ndef forward_diff(input_image, dim):\n    assert dim in (0, 1)\n    r, n = xp.shape(input_image)\n    size = xp.array((r, n))\n    position = xp.zeros(2, dtype=int)\n    temp1 = xp.zeros((r+1, n+1), dtype=float)\n    temp2 = xp.zeros((r+1, n+1), dtype=float)\n        \n    size[dim] += 1\n    position[dim] += 1\n\n    temp1[position[0]:size[0], position[1]:size[1]] = input_image\n    temp2[position[0]:size[0], position[1]:size[1]] = input_image\n    \n    size[dim] -= 1\n    temp2[0:size[0], 0:size[1]] = input_image\n    temp1 -= temp2\n    size[dim] += 1\n    return -temp1[position[0]:size[0], position[1]:size[1]]\n\ndef iter_deriv(input_image, b, scale, mu, dim1, dim2):\n    g = back_diff(forward_diff(input_image, dim1), dim2)\n    d = soft_threshold(g + b, 1 / mu)\n    b = b + (g - d)\n    L = scale * back_diff(forward_diff(d - b, dim2), dim1)\n    return L, b\n\ndef iter_xx(*args):\n    return iter_deriv(*args, dim1=1, dim2=1)\n\ndef iter_yy(*args):\n    return iter_deriv(*args, dim1=0, dim2=0)\n\ndef iter_xy(*args):\n    return iter_deriv(*args, dim1=0, dim2=1)\n\ndef iter_sparse(input_image, bsparse, scale, mu):\n    d = soft_threshold(input_image + bsparse, 1 / mu)\n    bsparse = bsparse + (input_image - d)\n    Lsparse = scale * (d - bsparse)\n    return Lsparse, bsparse\n\ndef denoise_image(input_image, iter_num=100, fidelity=150, sparsity_scale=10, continuity_scale=0.5, mu=1):\n    image_size = xp.shape(input_image)\n    #print(\"Initialize denoising\")\n    norm_array = (\n        operate_derivative(image_size, \"xx\") + \n        operate_derivative(image_size, \"yy\") + \n        2 * operate_derivative(image_size, \"xy\")\n    )\n    norm_array += (fidelity / mu) + sparsity_scale ** 2\n    b_arrays = {\n        \"xx\": xp.zeros(image_size, dtype=float),\n        \"yy\": xp.zeros(image_size, dtype=float),\n        \"xy\": xp.zeros(image_size, dtype=float),\n        \"L1\": xp.zeros(image_size, dtype=float),\n    }\n    g_update = xp.multiply(fidelity / mu, input_image)\n    for i in tqdm(range(iter_num), total=iter_num):\n        #print(f\"Starting iteration {i+1}\")\n        g_update = xp.fft.fftn(g_update)\n        if i == 0:\n            g = xp.fft.ifftn(g_update / (fidelity / mu)).real\n        else:\n            g = xp.fft.ifftn(xp.divide(g_update, norm_array)).real\n        g_update = xp.multiply((fidelity / mu), input_image)\n        \n        #print(\"XX update\")\n        L, b_arrays[\"xx\"] = iter_xx(g, b_arrays[\"xx\"], continuity_scale, mu)\n        g_update += L\n        \n        #print(\"YY update\")\n        L, b_arrays[\"yy\"] = iter_yy(g, b_arrays[\"yy\"], continuity_scale, mu)\n        g_update += L\n        \n        #print(\"XY update\")\n        L, b_arrays[\"xy\"] = iter_xy(g, b_arrays[\"xy\"], 2 * continuity_scale, mu)\n        g_update += L\n        \n        #print(\"L1 update\")\n        L, b_arrays[\"L1\"] = iter_sparse(g, b_arrays[\"L1\"], sparsity_scale, mu)\n        g_update += L\n        \n    g_update = xp.fft.fftn(g_update)\n    g = xp.fft.ifftn(xp.divide(g_update, norm_array)).real\n    \n    g[g < 0] = 0\n    g -= g.min()\n    g /= g.max()\n    return g\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\ndef rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    # pixels = (pixels >= thr).astype(int)\n    \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)\nclass CustomDataset(Dataset):\n    def __init__(self, images,  labels=None, transform=None):\n        self.images = images\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        image = self.images[idx]\n        image = torch.from_numpy(image.astype(np.float32)).unsqueeze(0).permute(0, 3, 1, 2)\n        return image\n    \n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_test_dataset(fragment_id):\n    test_images = read_image(fragment_id)\n    count=0\n    x1_list = list(range(0, test_images.shape[1]-tile_size+1, stride))\n    y1_list = list(range(0, test_images.shape[0]-tile_size+1, stride))\n    \n    test_images_list = []\n    xyxys = []\n    for y1 in y1_list:\n        for x1 in x1_list:\n            y2 = y1 + tile_size\n            x2 = x1 + tile_size\n            if np.all(test_images[y1:y2, x1:x2]==0):\n                count+=1    \n                continue\n            test_images_list.append(test_images[y1:y2, x1:x2])\n            xyxys.append((x1, y1, x2, y2))\n    xyxys = np.stack(xyxys)\n            \n    test_dataset = CustomDataset(test_images_list)\n    test_loader = DataLoader(test_dataset,\n                          batch_size=BATCH_SIZE,\n                          shuffle=False,\n                          num_workers=2, pin_memory=True, drop_last=False)\n    return test_loader, xyxys\n\ndef read_image(fragment_id):\n    images = []\n\n    idxs = range(Z_START,Z_START+Z_DIMS)\n    for i in tqdm(idxs):\n        image = cv2.imread(f\"{comp_dataset_path}/{mode}/{fragment_id}/surface_volume/{i:02}.tif\",0)\n        pad0 = (tile_size - image.shape[0] % tile_size)\n        pad1 = (tile_size - image.shape[1] % tile_size)\n\n        image = np.pad(image, [(0, pad0), (0, pad1)], constant_values=0)\n\n        images.append(image)\n    images = np.stack(images, axis=2)\n    return images","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Decoder(nn.Module):\n    def __init__(self, encoder_dims, upscale):\n        super().__init__()\n        self.convs = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(encoder_dims[i]+encoder_dims[i-1], encoder_dims[i-1], 3, 1, 1, bias=False),\n                nn.BatchNorm2d(encoder_dims[i-1]),\n                nn.ReLU(inplace=True)\n            ) for i in range(1, len(encoder_dims))])\n\n        self.logit = nn.Conv2d(encoder_dims[0], 1, 1, 1, 0)\n        self.up = nn.Upsample(scale_factor=upscale, mode=\"bilinear\")\n\n    def forward(self, feature_maps):\n        for i in range(len(feature_maps)-1, 0, -1):\n            f_up = F.interpolate(feature_maps[i], scale_factor=2, mode=\"bilinear\")\n            f = torch.cat([feature_maps[i-1], f_up], dim=1)\n            f_down = self.convs[i-1](f)\n            feature_maps[i-1] = f_down\n\n        x = self.logit(feature_maps[0])\n        mask = self.up(x)\n        return mask\nclass SegModel(nn.Module):\n    def __init__(self,model_depth):\n        super().__init__()\n        self.encoder = generate_model(model_depth=model_depth, n_input_channels=1)\n        self.decoder = Decoder(encoder_dims=[64, 128, 256, 512], upscale=4)\n        \n    def forward(self, x):\n        feat_maps = self.encoder(x)\n        feat_maps_pooled = [torch.mean(f, dim=2) for f in feat_maps]\n        pred_mask = self.decoder(feat_maps_pooled)\n        return pred_mask","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def TTA(x:torch.Tensor,model:nn.Module,tta=True,num_tta=4):\n    #x.shape=(batch,1,c,h,w)\n    if tta:\n        out=[]\n        shape=x.shape\n        xf1=torch.flip(x, dims=[3])#\n        xf2=torch.flip(x, dims=[4])#\n        x=[x,*[torch.rot90(x,k=i,dims=(-2,-1)) for i in range(1,num_tta)]]\n        for x_ in x:\n            out.append(model(x_)) \n        out.append(model(xf1))# \n        out.append(model(xf2)) #\n        del x_,xf1,xf2\n        gc.collect()\n        x=torch.cat(out,dim=0)\n        del out\n        gc.collect()\n        x=torch.sigmoid(x).squeeze(1)\n        x=x.reshape(num_tta+2,shape[0],*shape[3:])\n        xf = x[4:,:,:,:]\n        x =  x[:4,:,:,:]\n        x=[torch.rot90(x[i],k=-i,dims=(-2,-1)) for i in range(num_tta)]\n        xf1= torch.flip(xf[0], dims=[1])#\n        xf2= torch.flip(xf[1], dims=[2])#\n        xf=torch.cat([xf1,xf2])\n        del xf1,xf2;gc.collect();\n        x=torch.stack(x,dim=0)\n        x=torch.cat([x,xf.unsqueeze(1)])\n        return x.mean(0)\n    else :\n        x=model(x)\n        x=torch.sigmoid(x)\n        return x  \nclass EnsembleModel:\n    def __init__(self, use_tta=True):\n        self.models = []\n        self.use_tta = use_tta\n\n    def __call__(self, x):\n        output = [TTA(x,model)\n                   for model in self.models]\n        output=torch.stack(output,dim=0).mean(0)\n        return output\n    \n    def add_model(self, model):\n        self.models.append(model)\n\ndef build_ensemble_model(chkps):\n    model = EnsembleModel()\n    for model_path in chkps:\n        if '18' in model_path:\n            print('r18',model_path)\n            _model=SegModel(18)\n        else:\n            print('r34',model_path)\n            _model = SegModel(34)\n        _model.to(device)\n        state = torch.load(model_path)['model']\n        _model.load_state_dict(state,strict=True)\n        _model.eval()\n        model.add_model(_model)\n    \n    return model\n\ndef crop_center(img,cropx,cropy):\n    y,x = img.shape\n    startx = x//2-(cropx//2)\n    starty = y//2-(cropy//2)\n    return img[starty:starty+cropy,startx:startx+cropx]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if mode == 'test':\n    fragment_ids = sorted(os.listdir(comp_dataset_path + mode))\nelse:\n    fragment_ids = [3]\nprint(fragment_ids)    ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = build_ensemble_model(chkps)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer():\n    results = []\n    for fragment_id in fragment_ids:\n        test_loader, xyxys = make_test_dataset(fragment_id)\n        binary_mask = cv2.imread(comp_dataset_path + f\"{mode}/{fragment_id}/mask.png\", 0)\n        binary_mask = (binary_mask / 255).astype(int)\n        ori_h = binary_mask.shape[0]\n        ori_w = binary_mask.shape[1]\n        pad0 = (tile_size - binary_mask.shape[0] % tile_size)\n        pad1 = (tile_size - binary_mask.shape[1] % tile_size)\n        binary_mask = np.pad(binary_mask, [(0, pad0), (0, pad1)], constant_values=0)\n        mask_pred = np.zeros(binary_mask.shape)\n        mask_count = np.zeros(binary_mask.shape)\n        for step, (images) in tqdm(enumerate(test_loader), total=len(test_loader)):\n            images = images.to(device)  \n            bs = images.size(0)\n            with torch.no_grad():\n                y_preds = model(images).cpu().numpy()\n            start_idx = step*BATCH_SIZE\n            end_idx = start_idx + bs\n            for i, (x1, y1, x2, y2) in enumerate(xyxys[start_idx:end_idx]): \n                mask_pred[y1:y2, x1:x2] += np.pad(crop_center(y_preds[i],size,size),(expand//2,)) \n                mask_count[y1:y2, x1:x2] += np.pad(np.ones((size, size)),(expand//2,)) \n        del test_loader,images      \n        gc.collect()\n        torch.cuda.empty_cache()\n        plt.imshow(mask_count)\n        plt.show()\n        print(f'mask_count_min: {mask_count.min()}')\n        mask_pred /= (mask_count+1e-7)\n        \n        mask_pred=xp.array(mask_pred)\n        mask_pred=denoise_image(mask_pred, iter_num=250)\n        mask_pred=mask_pred.get()\n        \n        mask_pred = mask_pred[:ori_h, :ori_w]\n        binary_mask = binary_mask[:ori_h, :ori_w]\n        mask_pred = (mask_pred >= TH).astype(int)\n        mask_pred *= binary_mask \n        plt.imshow(mask_pred)\n        plt.show()    \n        inklabels_rle = rle(mask_pred) \n        results.append((fragment_id, inklabels_rle))    \n        del mask_pred, mask_count,y_preds,binary_mask,inklabels_rle\n        gc.collect()\n        torch.cuda.empty_cache()\n    return results","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_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)\n    print('Quiq save')\n    print(df_sub)\nelse:\n    results=infer()\n    sub = pd.DataFrame(results, columns=['Id', 'Predicted'])\n    sample_sub = pd.read_csv(comp_dataset_path + 'sample_submission.csv')\n    sample_sub = pd.merge(sample_sub[['Id']], sub, on='Id', how='left')\n    sample_sub.to_csv(\"submission.csv\", index=False)\n    sample_sub\n    ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}