{"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":"## Imports","metadata":{}},{"cell_type":"code","source":"!conda install ../input/pyvips-offline/*.tar.bz2 ","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T07:59:00.701478Z","iopub.execute_input":"2022-10-04T07:59:00.701801Z","iopub.status.idle":"2022-10-04T08:01:48.057431Z","shell.execute_reply.started":"2022-10-04T07:59:00.701723Z","shell.execute_reply":"2022-10-04T08:01:48.056203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport json\nimport glob\nimport pyvips\nimport zipfile\nimport warnings\nimport rasterio\nimport itertools\nimport skimage.io\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nimport plotly.graph_objects as go\n\nfrom tqdm.notebook import tqdm\nfrom skimage.io import imshow\nfrom scipy.ndimage import gaussian_filter\nfrom skimage.measure import label, regionprops, regionprops_table\n\nwarnings.filterwarnings(\"ignore\", category=rasterio.errors.NotGeoreferencedWarning)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:48.060001Z","iopub.execute_input":"2022-10-04T08:01:48.060387Z","iopub.status.idle":"2022-10-04T08:01:49.312684Z","shell.execute_reply.started":"2022-10-04T08:01:48.060338Z","shell.execute_reply":"2022-10-04T08:01:49.311764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Functions","metadata":{}},{"cell_type":"markdown","source":"### Params","metadata":{}},{"cell_type":"code","source":"# Hardcoded stuff, paths are to adapt to your setup\n\nimport torch\nimport numpy as np\n\nNUM_WORKERS = 2\n\nDATA_PATH = \"../input/mayo-clinic-strip-ai/\"\n\nLOG_PATH = \"../logs/\"\nOUT_PATH = \"/tmp/\"\n\nCLASSES = [\"CE\", \"LAA\"]\nNUM_CLASSES = len(CLASSES)\n\nMEAN = np.array([0.66437738, 0.50478148, 0.70114894])\nSTD = np.array([0.15825711, 0.24371008, 0.13832686])\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:49.314249Z","iopub.execute_input":"2022-10-04T08:01:49.314594Z","iopub.status.idle":"2022-10-04T08:01:50.616644Z","shell.execute_reply.started":"2022-10-04T08:01:49.314558Z","shell.execute_reply":"2022-10-04T08:01:50.615673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Utils","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nfrom itertools import product\n# from params import IMG_SIZE\n\n\ndef show_tiles(tiles, tiles_count, title=\"\"):\n    to_display = np.ones((tiles_count.v_tiles, tiles_count.h_tiles))\n    for i, tile in enumerate(tiles):\n        if tile is None:\n            to_display[i // tiles_count.h_tiles, i % tiles_count.h_tiles] = 0\n\n    cols_to_skip = np.argwhere(to_display.sum(0) == 0).flatten()\n    rows_to_skip = np.argwhere(to_display.sum(1) == 0).flatten()\n\n    figure, axes = plt.subplots(\n        tiles_count.v_tiles - len(rows_to_skip),\n        tiles_count.h_tiles - len(cols_to_skip),\n        figsize=(\n            (tiles_count.h_tiles - len(cols_to_skip)) * 3,\n            (tiles_count.v_tiles - len(rows_to_skip)) * 3,\n        ),\n    )\n    figure.suptitle(title, size=20, y=0.9)\n    axes = np.ravel(axes)\n\n    subplot_idx = 0\n    for i, tile in enumerate(tiles):\n        row, col = i // tiles_count.h_tiles, i % tiles_count.h_tiles\n\n        if row in rows_to_skip or col in cols_to_skip:\n            continue\n\n        if tile is not None:\n            axes[subplot_idx].imshow(tile)\n        else:\n            axes[subplot_idx].imshow(np.ones((IMG_SIZE, IMG_SIZE, 3)) * 0.98)\n        axes[subplot_idx].axis(\"off\")\n\n        subplot_idx += 1\n\n    figure.subplots_adjust(wspace=0, hspace=0.05)\n    figure.show()\n\n\ndef plot_matrix(mat, cmap=\"viridis\"):\n    \"\"\"\n    Plots a matrix.\n\n    Args:\n        mat (np array [n x n]): Matrix.\n        cmap (str, optional): Colormap name. Defaults to \"viridis\".\n    \"\"\"\n    n = mat.shape[0]\n    im_ = plt.imshow(mat, interpolation=\"nearest\", cmap=cmap)\n\n    # Display values\n    cmap_min, cmap_max = im_.cmap(0), im_.cmap(256)\n    thresh = (mat.max() + mat.min()) / 2.0\n    for i, j in product(range(n), range(n)):\n        color = cmap_max if mat[i, j] < thresh else cmap_min\n        text = f\"{mat[i, j]:.2f}\"\n        plt.text(j, i, text, ha=\"center\", va=\"center\", color=color)\n\n    plt.xticks(np.arange(n), np.arange(n))\n    plt.yticks(np.arange(n), np.arange(n))\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:50.620266Z","iopub.execute_input":"2022-10-04T08:01:50.620811Z","iopub.status.idle":"2022-10-04T08:01:50.635338Z","shell.execute_reply.started":"2022-10-04T08:01:50.62078Z","shell.execute_reply":"2022-10-04T08:01:50.634201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport torch\nimport random\nimport numpy as np\n\n\ndef seed_everything(seed):\n    \"\"\"\n    Seeds basic parameters for reproductibility of results.\n\n    Args:\n        seed (int): Number of the seed.\n    \"\"\"\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\n\ndef save_model_weights(model, filename, verbose=1, cp_folder=\"\"):\n    \"\"\"\n    Saves the weights of a PyTorch model.\n\n    Args:\n        model (torch model): Model to save the weights of.\n        filename (str): Name of the checkpoint.\n        verbose (int, optional): Whether to display infos. Defaults to 1.\n        cp_folder (str, optional): Folder to save to. Defaults to \"\".\n    \"\"\"\n\n    if verbose:\n        print(f\"\\n -> Saving weights to {os.path.join(cp_folder, filename)}\\n\")\n    torch.save(model.state_dict(), os.path.join(cp_folder, filename))\n\n\ndef load_model_weights(model, filename, verbose=1, cp_folder=\"\", strict=True):\n    \"\"\"\n    Loads the weights of a PyTorch model. The exception handles cpu/gpu incompatibilities.\n\n    Args:\n        model (torch model): Model to load the weights to.\n        filename (str): Name of the checkpoint.\n        verbose (int, optional): Whether to display infos. Defaults to 1.\n        cp_folder (str, optional): Folder to load from. Defaults to \"\".\n\n    Returns:\n        torch model: Model with loaded weights.\n    \"\"\"\n    state_dict = torch.load(os.path.join(cp_folder, filename), map_location=\"cpu\")\n\n    try:\n        model.load_state_dict(state_dict, strict=strict)\n        if verbose:\n            print(f\"\\n -> Loading weights from {os.path.join(cp_folder,filename)}\\n\")\n\n    except BaseException:\n        del state_dict['logits.weight'], state_dict['logits.bias']\n        model.load_state_dict(state_dict, strict=strict)\n\n        if verbose:\n            print(f\"\\n -> Loading encoder weights from {os.path.join(cp_folder,filename)}\\n\")\n\n    return model\n\n\ndef count_parameters(model, all=False):\n    \"\"\"\n    Count the parameters of a model.\n\n    Args:\n        model (torch model): Model to count the parameters of.\n        all (bool, optional):  Whether to count not trainable parameters. Defaults to False.\n\n    Returns:\n        int: Number of parameters.\n    \"\"\"\n\n    if all:\n        return sum(p.numel() for p in model.parameters())\n    else:\n        return sum(p.numel() for p in model.parameters() if p.requires_grad)\n\n\ndef worker_init_fn(worker_id):\n    \"\"\"\n    Handles PyTorch x Numpy seeding issues.\n\n    Args:\n        worker_id (int]): Id of the worker.\n    \"\"\"\n    np.random.seed(np.random.get_state()[1][0] + worker_id)\n\n    \nclass Config:\n    \"\"\"\n    Placeholder to load a config from a saved json\n    \"\"\"\n    def __init__(self, dic):\n        for k, v in dic.items():\n            setattr(self, k, v)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:50.637123Z","iopub.execute_input":"2022-10-04T08:01:50.637544Z","iopub.status.idle":"2022-10-04T08:01:50.654229Z","shell.execute_reply.started":"2022-10-04T08:01:50.637507Z","shell.execute_reply":"2022-10-04T08:01:50.653116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Preparaton","metadata":{}},{"cell_type":"code","source":"import cv2\nimport pyvips\nimport numpy as np\n\n\nformat_to_dtype = {\n    \"uchar\": np.uint8,\n    \"char\": np.int8,\n    \"ushort\": np.uint16,\n    \"short\": np.int16,\n    \"uint\": np.uint32,\n    \"int\": np.int32,\n    \"float\": np.float32,\n    \"double\": np.float64,\n    \"complex\": np.complex64,\n    \"dpcomplex\": np.complex128,\n}\n\n\ndef read_image_pyvips(path, max_size=20000):\n    image = pyvips.Image.thumbnail(path, max_size)\n\n    image = np.ndarray(\n        buffer=image.write_to_memory(),\n        dtype=format_to_dtype[image.format],\n        shape=[image.height, image.width, image.bands],\n    )\n\n    return image\n\n\ndef get_img_prop(tile, min_saturation=20):\n    hsv = cv2.cvtColor(tile, cv2.COLOR_RGB2HSV)\n    _, s, _ = cv2.split(hsv)\n\n    low_sat = (s < min_saturation).mean()\n\n    counts = np.bincount((tile.mean(-1).flatten().astype(int)))\n    background_prop = counts[np.argsort(counts)[::-1][:3]].sum() / counts.sum()\n\n    background_prop = max(background_prop, low_sat)\n\n    return 1 - background_prop\n\n\ndef get_grid(orig_size, tile_size, overlap_factor=1):\n    top_x = np.arange(\n        orig_size[0] % tile_size // 2,  # shift to center grid\n        orig_size[0],\n        int(tile_size / overlap_factor),\n    )[:-1]\n    top_y = np.arange(\n        orig_size[1] % tile_size // 2,  # shift to center grid\n        orig_size[1],\n        int(tile_size / overlap_factor),\n    )[:-1]\n    grid = []\n    for x in top_x:\n        right_space = orig_size[0] - (x + tile_size)\n        if right_space > 0:\n            boundaries_x = (x, x + tile_size)\n        else:\n            boundaries_x = (x + right_space, x + right_space + tile_size)\n\n        for y in top_y:\n            down_space = orig_size[1] - (y + tile_size)\n            if down_space > 0:\n                boundaries_y = (y, y + tile_size)\n            else:\n                boundaries_y = (y + down_space, y + down_space + tile_size)\n            grid.append((boundaries_x, boundaries_y))\n\n    return grid\n\n\ndef remove_chunks(image, size=256, min_prop=0.01):\n    h, w, _ = image.shape\n\n    top_x = np.arange(0, h, size)\n    top_y = np.arange(0, w, size)\n\n    kept = []\n    for x in top_x:\n        prop = get_img_prop(image[x: x + size])\n        if prop > min_prop:\n            kept.append(image[x: x + size])\n\n    #         plt.figure(figsize=(15, 3))\n    #         plt.imshow(image[x: x+size])\n    #         plt.title(prop)\n    #         plt.show()\n\n    if len(kept):\n        image = np.concatenate(kept, 0)\n\n    kept = []\n    for y in top_y:\n        prop = get_img_prop(image[:, y: y + size])\n        if prop > min_prop:\n            kept.append(image[:, y: y + size])\n\n    if len(kept):\n        image = np.concatenate(kept, 1)\n\n    return image\n\n\ndef get_normalization_ratio(img, tile_size=512, min_saturation=20):\n    h, w, _ = img.shape\n    grid = get_grid((h, w), tile_size)\n\n    for x, y in grid:\n        tile = img[x[0]: x[1], y[0]: y[1]]\n        img_prop = get_img_prop(tile, min_saturation=min_saturation)\n\n        if img_prop < 0.25:\n            ratio = np.mean(tile, (0, 1)) / np.array([255.0, 255.0, 255.0])\n            return ratio\n\n    return None\n\n\ndef normalize(image, ratio=None):\n    if ratio is None:\n        return image\n    else:\n        return np.clip((image / ratio), 0, 255).astype(np.uint8)\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:50.655871Z","iopub.execute_input":"2022-10-04T08:01:50.656362Z","iopub.status.idle":"2022-10-04T08:01:50.678627Z","shell.execute_reply.started":"2022-10-04T08:01:50.656326Z","shell.execute_reply":"2022-10-04T08:01:50.677629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nfrom collections import Counter\nfrom scipy.ndimage import gaussian_filter\nfrom skimage.morphology import erosion, dilation\n\n\ndef multi_dil(im, num):\n    element = np.array([[0, 1, 0], [1, 1, 1], [0, 1, 0]])\n    for i in range(num):\n        im = dilation(im, element)\n    return im\n\n\ndef multi_ero(im, num):\n    element = np.array([[0, 1, 0], [1, 1, 1], [0, 1, 0]])\n    for i in range(num):\n        im = erosion(im, element)\n    return im\n\n\ndef denoise(binarized, n=10, m=2):\n    # Remove small comps\n    binarized = multi_ero(binarized, m)\n    binarized = multi_dil(binarized, m)\n\n    # merge big comps\n    binarized = multi_dil(binarized, n)\n    binarized = multi_ero(binarized, n)\n\n    # return binarized\n\n    # Remove small comps\n    binarized = multi_ero(binarized, m)\n    return multi_dil(binarized, m)\n\n\ndef remove_bubble(img):\n    bubble = (img.mean(-1) < 250) & (img.mean(-1) > 220) & (img.std(-1) < 5)\n    bubble = bubble.astype(float)\n\n    bubble = multi_dil(bubble, 2)\n    bubble = multi_ero(bubble, 5)\n    bubble = multi_dil(bubble, 3)\n\n    # plt.imshow(bubble)\n    bubble = bubble[..., None].astype(int)\n\n    return img * (1 - bubble) + bubble * 255\n\n\ndef remove_background(grayscale, img=None):\n    counts = np.bincount(grayscale.flatten())\n\n    bg_prop = np.max(counts) / grayscale.size\n    bg = np.argmax(counts)\n\n    if img is None:\n        if bg_prop > 0.2:\n            grayscale[grayscale == bg] = 255\n        return grayscale\n    else:\n        if bg_prop > 0.2:\n            img[grayscale[..., None] == bg] = 255\n        return img\n\n\ndef euclidian_dist(image1, image2):\n    image1 = cv2.resize(image1, (256, 256)).astype(float) / 255\n    image2 = cv2.resize(image2, (256, 256)).astype(float) / 255\n\n    image1 = image1.mean(-1)\n    image2 = image2.mean(-1)\n\n    image1 = gaussian_filter(image1, sigma=1)\n    image2 = gaussian_filter(image2, sigma=1)\n\n    return ((image1 < 0.98) != (image2 < 0.98)).astype(float).mean() ** 2\n    # return (np.abs(image1 - image2)).mean()\n\n\ndef shape_dist(image1, image2):\n    h1, w1 = image1.shape[:2]\n    h2, w2 = image2.shape[:2]\n    return abs(h1 - h2) / min(h1, h2) + abs(w1 - w2) / min(w1, w2)\n\n\ndef component_distance(img, labeled, df, plot=False, min_size=500):\n    crops, to_ignore = [], []\n\n    min_size = min(df[\"area\"].max(), min_size)\n\n    for i in range(len(df)):\n        x0, y0, x1, y1 = df[[\"bbox-0\", \"bbox-1\", \"bbox-2\", \"bbox-3\"]].values[i]\n\n        if df[\"area\"].values[i] < min_size:\n            to_ignore.append(i)\n\n        crop = img[x0:x1, y0:y1].astype(np.uint8)\n        # crop = crop\n\n        if plot:\n            plt.imshow(crop)\n            plt.title(i)\n            plt.show()\n\n        crops.append(crop)\n\n    #     dists = np.zeros((len(df), len(df)))\n    dists = np.eye(len(df))\n    for i in range(len(df)):\n        for j in range(i):\n            #             d = hash_dist(crops[i], crops[j])\n            d = euclidian_dist(crops[i], crops[j])\n            d2 = shape_dist(crops[i], crops[j])\n\n            dists[i, j] = d + 0.5 * (d2 > 0.4)\n            dists[j, i] = d + 0.5 * (d2 > 0.4)\n\n    dists[to_ignore] = 1\n    dists[:, to_ignore] = 1\n\n    return crops, dists, to_ignore\n\n\ndef get_duplicates(dists, threshold=0.15):\n    duplicates = []\n\n    for i in range(len(dists)):\n        # min_dist_i = np.min(dists[i])\n\n        for j in range(i):\n            # min_dist_j = np.min(dists[j])\n            dist = dists[i, j]\n            # print(min_score, score)\n\n            if (\n                dist < threshold\n                # and\n                # ((dist <= min_dist_i and dist <= min_dist_j) or dist <= threshold / 2)\n            ):\n                duplicates.append((i, j, dist))\n    duplicates = sorted(duplicates, key=lambda x: x[2])\n    return duplicates\n\n\ndef get_clusters(df_comp, duplicates, max_size=3):\n    df_comp['cluster'] = -1\n\n    for i, d in enumerate(duplicates):\n        clust_idx = i\n\n        if df_comp['cluster'][d[0]] >= 0 and df_comp['cluster'][d[1]] >= 0:\n            continue\n\n        elif df_comp['cluster'][d[0]] >= 0:\n            clust_idx = df_comp['cluster'][d[0]]\n            if Counter(df_comp['cluster'])[clust_idx] >= max_size:\n                continue\n                # clust_idx = i\n\n        elif df_comp['cluster'][d[1]] >= 0:\n            clust_idx = df_comp['cluster'][d[1]]\n            if Counter(df_comp['cluster'])[clust_idx] >= max_size:\n                continue\n\n        df_comp['cluster'][d[0]] = clust_idx\n        df_comp['cluster'][d[1]] = clust_idx\n\n\ndef get_roi(df_comp, shape, found_crop=False, min_size=500, margin=5):\n\n    order = ['bbox-0', 'bbox-1'] if shape[0] >= shape[1] else ['bbox-1', 'bbox-0']\n    order = ['-area'] + order if found_crop else order\n\n    df_comp = df_comp.sort_values(order).reset_index(drop=True)\n\n    m = margin\n    df_comp = df_comp[df_comp['area'] > min_size]\n\n    if df_comp['cluster'].max() >= 0:  # has clusters, keep one per cluster.\n        if found_crop or df_comp['cluster'].max() >= 2:  # one per unclustered + increase margin\n            df_comp = df_comp[(~df_comp.duplicated(subset=\"cluster\"))]\n            m = margin * 5\n        else:\n            df_comp = df_comp[(df_comp['cluster'] == -1) | (~df_comp.duplicated(subset=\"cluster\"))]\n\n        # if -1 in df_comp['cluster'].values:  # unmatched elements, extend margin\n        #     m = margin * 5\n        # df_comp = df_comp[df_comp[\"cluster\"] != -1]\n\n    x0, y0 = df_comp[['bbox-0', 'bbox-1']].min()\n    x1, y1 = df_comp[['bbox-2', 'bbox-3']].max()\n\n    x0 = max(0, x0 - m)\n    y0 = max(0, y0 - m)\n    x1 = min(shape[0], x1 + m)\n    y1 = min(shape[1], y1 + m)\n\n    return x0, x1, y0, y1, df_comp\n\n\ndef random_crop(crop1, crop2):\n    \"\"\"\n    Crop crop1 to size of crop2.\n\n    Args:\n        crop1 (_type_): _description_\n        crop2 (_type_): _description_\n\n    Returns:\n        _type_: _description_\n    \"\"\"\n    shape1 = crop1.shape[:2]\n    shape2 = crop2.shape[:2]\n\n    crop = crop1.copy()\n\n    if shape1[0] > shape2[0]:\n        x_start = np.random.randint(shape1[0] - shape2[0])\n        crop = crop[x_start: x_start + shape2[0]]\n\n    if shape1[1] > shape2[1]:\n        y_start = np.random.randint(shape1[1] - shape2[1])\n        crop = crop[:, y_start: y_start + shape2[1]]\n\n    return crop\n\n\ndef match_crops(crop1, crop2, threshold=0.05, trials=100, rot=False, plot=False):\n\n    for t in range(trials):\n        if 0 < t < 4:  # try rotations\n            th = threshold\n            cropr = np.rot90(crop1.copy(), k=t)\n        else:  # crop\n            th = threshold * 2/3\n            cropr = random_crop(crop1, crop2)\n\n        if rot:\n            if np.random.random() < 0.5:\n                cropr = np.rot90(cropr, k=np.random.randint(1, 4))\n\n        dist = euclidian_dist(crop2.copy(), cropr.copy())\n\n        if dist < th:\n            if plot:\n                plt.figure(figsize=(10, 5))\n                plt.subplot(1, 2, 1)\n                plt.imshow(crop2)\n                plt.title(\"orig\")\n\n                plt.subplot(1, 2, 2)\n                plt.imshow(cropr)\n                plt.title(f\"crop - d={dist :.3f}\")\n                plt.show()\n\n            return True\n\n    return False\n\n\ndef compute_dist_matrix(df):\n    mat = np.ones((len(df), len(df))) * 100\n    for i in range(len(df)):\n        for j in range(i):\n            dx = np.abs(df[\"xc\"][i] - df[\"xc\"][j]) / (0.5 * df[\"h\"][i] + 0.5 * df[\"h\"][j])\n            dy = np.abs(df[\"yc\"][i] - df[\"yc\"][j]) / (0.5 * df[\"w\"][i] + 0.5 * df[\"w\"][j])\n            d = np.sqrt(dx ** 2 + dy ** 2)\n            mat[i, j] = d\n            mat[j, i] = d\n    return mat\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:50.680322Z","iopub.execute_input":"2022-10-04T08:01:50.681065Z","iopub.status.idle":"2022-10-04T08:01:50.779294Z","shell.execute_reply.started":"2022-10-04T08:01:50.681023Z","shell.execute_reply":"2022-10-04T08:01:50.778468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset","metadata":{}},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport albumentations as albu\nfrom albumentations.pytorch import ToTensorV2\n# from params import MEAN, STD\n\n\ndef pad_to_square(img, pad_value=255):\n    pad = abs(img.shape[0] - img.shape[1]) / 2\n    pad_1, pad_2 = int(np.floor(pad)), int(np.ceil(pad))\n\n    if img.shape[0] > img.shape[1]:\n        axis = 1\n        shape_1 = (img.shape[0], pad_1, img.shape[2])\n        shape_2 = (img.shape[0], pad_2, img.shape[2])\n    elif img.shape[0] < img.shape[1]:\n        axis = 0\n        shape_1 = (pad_1, img.shape[1], img.shape[2])\n        shape_2 = (pad_2, img.shape[1], img.shape[2])\n    else:\n        return img\n\n    padding_1 = pad_value * np.ones(shape_1, dtype=img.dtype)\n    padding_2 = pad_value * np.ones(shape_2, dtype=img.dtype)\n\n    return np.concatenate([padding_1, img, padding_2], axis=axis)\n\n\ndef blur_transforms(p=0.5, blur_limit=3, gaussian_limit=(3, 5)):\n    \"\"\"\n    Applies MotionBlur or GaussianBlur random with a probability p.\n\n    Args:\n        p (float, optional): probability. Defaults to 0.5.\n        blur_limit (int, optional): Blur intensity limit. Defaults to 5.\n\n    Returns:\n        albumentation transforms: transforms.\n    \"\"\"\n    return albu.OneOf(\n        [\n            albu.MotionBlur(blur_limit=blur_limit, always_apply=True),\n            albu.GaussianBlur(blur_limit=gaussian_limit, always_apply=True),\n        ],\n        p=p,\n    )\n\n\ndef color_transforms(p=0.5):\n    \"\"\"\n    Applies RandomGamma or RandomBrightnessContrast random with a probability p.\n\n    Args:\n        p (float, optional): probability. Defaults to 0.5.\n\n    Returns:\n        albumentation transforms: transforms.\n    \"\"\"\n    return albu.OneOf(\n        [\n            albu.Compose(\n                [\n                    albu.RandomGamma(gamma_limit=(80, 120), p=1),\n                    albu.RandomBrightnessContrast(\n                        brightness_limit=(0, 0.1),\n                        contrast_limit=0.1,\n                        p=1,\n                    ),\n                ]\n            ),\n            albu.RGBShift(\n                r_shift_limit=20,\n                g_shift_limit=20,\n                b_shift_limit=20,\n                p=1,\n            ),\n            albu.HueSaturationValue(\n                hue_shift_limit=20,\n                sat_shift_limit=20,\n                val_shift_limit=20,\n                p=1,\n            ),\n            albu.ColorJitter(\n                brightness=(1, 1.3),\n                contrast=(0.9, 1.3),\n                saturation=(0.9, 1.3),\n                hue=0.1,\n                p=1,\n            ),\n        ],\n        p=p,\n    )\n\n\ndef deformation_transform(p=0.5):\n    \"\"\"\n    Applies ElasticTransform, GridDistortion or OpticalDistortion with a probability p.\n\n    Args:\n        p (float, optional): probability. Defaults to 0.5.\n\n    Returns:\n        albumentation transforms: transforms.\n    \"\"\"\n    return albu.OneOf(\n        [\n            albu.ElasticTransform(\n                alpha=1,\n                sigma=25,\n                alpha_affine=25,\n                always_apply=True,\n            ),\n            albu.GridDistortion(always_apply=True),\n        ],\n        p=p,\n    )\n\n\ndef center_crop(size):\n    \"\"\"\n    Applies a padded center crop.\n\n    Args:\n        size (int): Crop size.\n\n    Returns:\n        albumentation transforms: transforms.\n    \"\"\"\n    if size is None:  # disable cropping\n        p = 0\n        size = 256\n    else:  # always crop\n        p = 1\n\n    return albu.Compose(\n        [\n            albu.PadIfNeeded(size, size, p=p, border_mode=cv2.BORDER_CONSTANT),\n            albu.CenterCrop(size, size, p=p),\n        ],\n        p=1,\n    )\n\n\ndef get_transfos(augment=True, visualize=False, mean=MEAN, std=STD, resize=None):\n    \"\"\"\n    Returns transformations.\n\n    Args:\n        augment (bool, optional): Whether to apply augmentations. Defaults to True.\n        visualize (bool, optional): Whether to use transforms for visualization. Defaults to False.\n        mean (np array, optional): Mean for normalization. Defaults to MEAN.\n        std (np array, optional): Standard deviation for normalization. Defaults to STD.\n\n    Returns:\n        albumentation transforms: transforms.\n    \"\"\"\n    resize_aug = [albu.Resize(resize, resize)] if resize is not None else []\n    if visualize:\n        normalizer = albu.Compose(\n            resize_aug\n            + [\n                albu.Normalize(mean=[0, 0, 0], std=[1, 1, 1]),\n                ToTensorV2(),\n            ],\n            p=1,\n        )\n    else:\n        normalizer = albu.Compose(\n            resize_aug\n            + [\n                albu.Normalize(mean=mean, std=std),\n                ToTensorV2(),\n            ],\n            p=1,\n        )\n\n    if augment:\n        return albu.Compose(\n            [\n                albu.VerticalFlip(p=0.5),\n                albu.HorizontalFlip(p=0.5),\n                albu.ShiftScaleRotate(\n                    scale_limit=0.1,\n                    shift_limit=0.,\n                    rotate_limit=45,\n                    p=0.5,\n                ),\n                color_transforms(p=0.5),\n                # deformation_transform(p=0.25),\n                # blur_transforms(p=0.25),\n                normalizer,\n            ]\n        )\n    else:\n        return normalizer","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:50.781599Z","iopub.execute_input":"2022-10-04T08:01:50.782197Z","iopub.status.idle":"2022-10-04T08:01:51.40463Z","shell.execute_reply.started":"2022-10-04T08:01:50.78216Z","shell.execute_reply":"2022-10-04T08:01:51.403673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport torch\nimport numpy as np\nimport pandas as pd\nfrom torch.utils.data import Dataset\n\n# from params import CLASSES\n# from data.transforms import pad_to_square\n\n\ndef filter_tiles(df, min_per_image=4, max_per_image=None, min_prop=0):\n    filtered_dfs = []\n    for img, df_img in df.groupby('image_id'):\n        df_img = df_img.sort_values('img_prop', ascending=False).reset_index(drop=True)\n        if max_per_image is not None:\n            df_img = df_img.head(max_per_image)\n        filtered_dfs.append(\n            df_img[(df_img['img_prop'] > min_prop) | (df_img.index < min_per_image)]\n        )\n\n    return pd.concat(filtered_dfs, ignore_index=True)\n\n\ndef prepare_data(img_folder, tile=False, min_prop=0, min_per_image=0):\n    df = pd.read_csv(img_folder + \"df.csv\")\n    df['target'] = df['label'].apply(lambda x: CLASSES.index(x))\n\n    if \"img_prop\" not in df.columns:\n        df[\"img_prop\"] = 1\n\n    if not tile:\n        return df\n\n    df = filter_tiles(df, min_prop=min_prop, min_per_image=min_per_image)\n\n    dfg = df.groupby('image_id').agg(list).reset_index()\n    for c in ['center_id', 'patient_id', 'image_num', 'label', 'target']:\n        dfg[c] = dfg[c].apply(lambda x: x[0])\n\n    return dfg\n\n\nclass TileDataset(Dataset):\n    \"\"\"\n    Segmentation dataset for training / validation on tiles.\n    \"\"\"\n    def __init__(self, df, n_tiles=0, transforms=None, train=False):\n        \"\"\"\n        Constructor.\n\n        Args:\n            df (pandas dataframe): Metadata.\n            transforms (albumentation transforms, optional): Transforms to apply. Defaults to None.\n        \"\"\"\n        self.df = df\n        self.n_tiles = n_tiles\n        self.transforms = transforms\n\n        self.paths = df[\"path\"].values\n        self.props = df[\"img_prop\"].values\n        self.targets = df[\"target\"].values\n\n        self.train = train\n        self.min_prop = 0.25\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    def _getimages(self, paths, props):\n        images = []\n        for path in paths:\n            image = cv2.imread(path)\n\n            image = pad_to_square(image)\n\n            if self.transforms:\n                transformed = self.transforms(image=image)\n                image = transformed[\"image\"]\n\n            images.append(image)\n\n        while len(images) < self.n_tiles:\n            images.append(torch.zeros(image.size(), dtype=image.dtype))\n            props = np.append(props, 0)\n\n        return torch.stack(images), torch.tensor(props).float()\n\n    def __getitem__(self, idx):\n\n        if self.n_tiles:\n            paths = np.array(self.paths[idx])\n            props = np.array(self.props[idx])\n\n            if self.train and len(paths) >= self.n_tiles:  # sort ?\n                indices = np.random.choice(np.arange(len(paths)), self.n_tiles, replace=False)\n            else:\n                indices = np.argsort(props)[::-1][:self.n_tiles]\n\n            paths = paths[indices]\n            props = props[indices]\n        else:\n            paths = [self.paths[idx]]\n            props = np.ones(1)\n\n        images, props = self._getimages(paths, props)\n\n        y = torch.tensor(self.targets[idx])\n\n        w = props > self.min_prop\n        w[0] = 1\n        w = w.float()\n\n        return images, y, w\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:51.406114Z","iopub.execute_input":"2022-10-04T08:01:51.406692Z","iopub.status.idle":"2022-10-04T08:01:51.425746Z","shell.execute_reply.started":"2022-10-04T08:01:51.406655Z","shell.execute_reply":"2022-10-04T08:01:51.424747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('../input/timm-0-6-9/pytorch-image-models-master')\nimport timm\nimport torch.nn as nn\n\n# from utils.torch import load_model_weights\n\n\ndef define_model(\n    name,\n    num_classes=1,\n    pretrained_weights=\"\",\n    pretrained=True,\n    average=\"ft\"\n):\n    \"\"\"\n    Loads a pretrained model & builds the architecture.\n    Supports timm models.\n\n    Args:\n        name (str): Model name\n        num_classes (int, optional): Number of classes. Defaults to 1.\n        pretrained_weights (str, optional): Path to pretrained encoder weights. Defaults to ''.\n        pretrained (bool, optional): Whether to load timm pretrained weights.\n\n    Returns:\n        torch model -- Pretrained model.\n    \"\"\"\n    # Load pretrained model\n    encoder = getattr(timm.models, name)(pretrained=pretrained)\n    encoder.name = name\n\n    # Tile Model\n    model = TileModel(\n        encoder,\n        num_classes=num_classes,\n        average=average,\n    )\n\n    if pretrained_weights:\n        # raise NotImplementedError\n        model = load_model_weights(model, pretrained_weights, verbose=1, strict=False)\n\n    return model\n\n\nclass TileModel(nn.Module):\n    \"\"\"\n    Model with an attention mechanism.\n    \"\"\"\n    def __init__(\n        self,\n        encoder,\n        num_classes=1,\n        average=\"ft\",\n    ):\n        \"\"\"\n        Constructor.\n\n        Args:\n            encoder (timm model): Encoder.\n            num_classes (int, optional): Number of classes. Defaults to 1.\n        \"\"\"\n        super().__init__()\n\n        self.encoder = encoder\n        self.nb_ft = encoder.num_features\n        self.num_classes = num_classes\n\n        assert average in [\"ft\", \"proba\"], \"Averaging not supported\"\n        self.average = average\n\n        self.logits = nn.Linear(self.nb_ft, num_classes)\n\n    def extract_features(self, x):\n        \"\"\"\n        Extract features function.\n\n        Args:\n            x (torch tensor [batch_size x 3 x w x h]): Input batch.\n\n        Returns:\n            torch tensor [batch_size x num_features]: Features.\n        \"\"\"\n        fts = self.encoder.forward_features(x)\n\n        if len(fts.size()) >= 4:  # cnn\n            return fts.mean(-1).mean(-1)\n        else:\n            return fts.mean(-2)  # transfo\n\n    def forward(self, x, w=None):\n        \"\"\"\n        Forward function.\n\n        Args:\n            x (torch tensor [batch_size x 3 x w x h]): Input batch.\n\n        Returns:\n            torch tensor [batch_size x num_classes]: logits.\n        \"\"\"\n        n_tiles = x.size(1) if len(x.size()) > 4 else 1\n        x = x.view(-1, *x.size()[-3:]).contiguous()  # bs x n_tiles x c x h x w -> bs*n_tiles x ...\n\n        fts = self.extract_features(x)\n        fts = fts.view(-1, n_tiles, self.nb_ft)  # bs*n_tiles x nb_ft -> bs x n_tiles x nb_ft\n\n        if self.average == \"ft\":\n            if w is None:\n                fts = fts.mean(1)  # avg pooling\n            else:\n                w = w.unsqueeze(-1)\n                fts = (fts * w).sum(1) / w.sum(1)  # masked pooling\n\n        logits = self.logits(fts)\n\n        if self.average == \"proba\":\n            if w is None:\n                logits = logits.mean(1)  # avg pooling\n            else:\n                w = w.unsqueeze(-1)\n                logits = (logits * w).sum(1) / w.sum(1)  # masked pooling\n\n        return logits\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:51.430276Z","iopub.execute_input":"2022-10-04T08:01:51.430666Z","iopub.status.idle":"2022-10-04T08:01:53.086066Z","shell.execute_reply.started":"2022-10-04T08:01:51.430635Z","shell.execute_reply":"2022-10-04T08:01:53.085019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inference","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nfrom torch.utils.data import DataLoader\n\n# from params import NUM_WORKERS\n\nFLIPS = [None, [-1], [-2], [-2, -1]]\n\n\ndef predict(model, dataset, loss_config, batch_size=64, device=\"cuda\"):\n    \"\"\"\n    Torch predict function.\n\n    Args:\n        model (torch model): Model to predict with.\n        dataset (CustomDataset): Dataset to predict on.\n        loss_config (dict): Loss config, used for activation functions.\n        batch_size (int, optional): Batch size. Defaults to 64.\n        device (str, optional): Device for torch. Defaults to \"cuda\".\n\n    Returns:\n        numpy array [len(dataset) x num_classes]: Predictions.\n    \"\"\"\n    model.eval()\n    preds = np.empty((0,  model.num_classes))\n\n    loader = DataLoader(\n        dataset, batch_size=batch_size, shuffle=False, num_workers=NUM_WORKERS\n    )\n\n    with torch.no_grad():\n        for batch in loader:\n            x = batch[0].to(device)\n\n            # Forward\n            pred = model(x)\n\n            # Get probabilities\n            if loss_config['activation'] == \"sigmoid\":\n                pred = pred.sigmoid()\n            elif loss_config['activation'] == \"softmax\":\n                pred = pred.softmax(-1)\n            preds = np.concatenate([preds, pred.cpu().numpy()])\n\n    return preds\n\n\ndef predict_tta(model, dataset, loss_config, batch_size=64, device=\"cuda\"):\n    \"\"\"\n    Torch predict function with flip TTA.\n\n    Args:\n        model (torch model): Model to predict with.\n        dataset (CustomDataset): Dataset to predict on.\n        loss_config (dict): Loss config, used for activation functions.\n        batch_size (int, optional): Batch size. Defaults to 64.\n        device (str, optional): Device for torch. Defaults to \"cuda\".\n\n    Returns:\n        numpy array [len(dataset) x num_classes]: Predictions.\n    \"\"\"\n    model.eval()\n    preds = np.empty((0,  model.num_classes))\n\n    loader = DataLoader(\n        dataset, batch_size=batch_size, shuffle=False, num_workers=NUM_WORKERS\n    )\n\n    with torch.no_grad():\n        for batch in loader:\n            x = batch[0].to(device)\n            preds_tta = []\n\n            for f in FLIPS:\n                # Forward\n                pred = model(torch.flip(x, f) if f is not None else x)\n\n                # Get probabilities\n                if loss_config['activation'] == \"sigmoid\":\n                    pred = pred.sigmoid()\n                elif loss_config['activation'] == \"softmax\":\n                    pred = pred.softmax(-1)\n                preds_tta.append(pred.cpu().numpy())\n\n            preds = np.concatenate([preds, np.mean(preds_tta, 0)])\n\n    return preds\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:53.090133Z","iopub.execute_input":"2022-10-04T08:01:53.090481Z","iopub.status.idle":"2022-10-04T08:01:53.104732Z","shell.execute_reply.started":"2022-10-04T08:01:53.090451Z","shell.execute_reply":"2022-10-04T08:01:53.103571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def scale(preds, min_=0.2, max_=0.8):  # def scale(preds, min_=0.05, max_=0.55):\n    preds = (preds - preds.min()) / (preds.max() - preds.min())\n    preds = preds * (max_ - min_) + min_\n\n    return preds\n\n\ndef shrink(x):\n    if x < 0.3:\n        return 0.2 + x / 3\n    elif x > 0.7:\n        return 0.8 - ((1 - x) / 3)\n    else:\n        return x\n","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T08:01:53.106304Z","iopub.execute_input":"2022-10-04T08:01:53.106945Z","iopub.status.idle":"2022-10-04T08:01:53.11767Z","shell.execute_reply.started":"2022-10-04T08:01:53.10691Z","shell.execute_reply":"2022-10-04T08:01:53.116672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preparation","metadata":{}},{"cell_type":"markdown","source":"#### Params","metadata":{}},{"cell_type":"code","source":"MIN_SIZE = 750 # 500\nN = 6\nM = 2\nMARGIN = 5\nTHRESHOLD_DUP = 0.1\nSIZE = 256\nBIN_THRESH = 240\nSIGMA = 2\nMIN_PROP = 0.02\n\nIMG_SIZE = 1024  # 512\n\nSAVE = True\n\nSAVE_FOLDER_D = OUT_PATH + f\"train_d_{IMG_SIZE}/\"\nSAVE_FOLDER = OUT_PATH + f\"train_{IMG_SIZE}/\"","metadata":{"execution":{"iopub.status.busy":"2022-10-04T08:01:53.119214Z","iopub.execute_input":"2022-10-04T08:01:53.119657Z","iopub.status.idle":"2022-10-04T08:01:53.12989Z","shell.execute_reply.started":"2022-10-04T08:01:53.119516Z","shell.execute_reply":"2022-10-04T08:01:53.128834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Data","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv(os.path.join(DATA_PATH, \"test.csv\"))\ntest_paths = os.listdir(os.path.join(DATA_PATH, \"test\"))\n\nif not os.path.exists(OUT_PATH):\n    os.mkdir(OUT_PATH)\n    \nif SAVE and not os.path.exists(SAVE_FOLDER):\n    os.mkdir(SAVE_FOLDER)\n\nif SAVE and not os.path.exists(SAVE_FOLDER_D):\n    os.mkdir(SAVE_FOLDER_D)\n    \nPLOT = len(test_df) < 10","metadata":{"execution":{"iopub.status.busy":"2022-10-04T08:01:53.131693Z","iopub.execute_input":"2022-10-04T08:01:53.132117Z","iopub.status.idle":"2022-10-04T08:01:53.152129Z","shell.execute_reply.started":"2022-10-04T08:01:53.132083Z","shell.execute_reply":"2022-10-04T08:01:53.151282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Main","metadata":{}},{"cell_type":"code","source":"metadata = []\npath = \"\"\nseed_everything(42)\n\nfor idx in tqdm(range(len(test_df))):\n    image_id = test_df[\"image_id\"][idx]\n\n#     image_id = \"0d93ce_0\"\n    img_path = f\"test/{image_id}.tif\"\n\n    image = read_image_pyvips(DATA_PATH + img_path, max_size=20000)\n    \n#     image = np.random.randint(0, 255, (20000, 20000, 3), dtype=np.uint8)\n\n    ##########\n    # SIMPLE # \n    ##########\n    ratio = get_normalization_ratio(image)\n\n    sz = np.array(image.shape[:2]) \n    sz = (sz * 0.8).astype(int)\n    img = cv2.resize(image, sz[::-1])\n    img = remove_chunks(img, 128, min_prop=MIN_PROP)\n\n#     img = remove_chunks(image, 256, min_prop=MIN_PROP)\n    img = normalize(img, ratio)\n    img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n\n    if PLOT:\n        plt.figure(figsize=(10, 10))\n        plt.imshow(img)\n        plt.axis(False)\n        plt.title(image_id)\n        plt.show()\n    \n    path = SAVE_FOLDER + f\"{image_id}.png\"\n    if SAVE:\n        cv2.imwrite(path, img)\n        \n    del img\n    gc.collect()\n    \n    ############\n    # ADVANCED #\n    ############\n\n    orig_shape = image.shape[:2]\n\n    sz = np.array(image.shape[:2]) \n    factor = np.min(sz) / SIZE\n    sz = (sz / factor).astype(int)\n    \n    viz = cv2.resize(image, sz[::-1])\n    viz = normalize(viz, ratio)\n    \n    viz = remove_bubble(viz)\n    \n    grayscale = viz.mean(-1).astype(int)\n    grayscale = remove_background(grayscale)\n    grayscale = gaussian_filter(grayscale, sigma=SIGMA)\n\n    binarized = (grayscale < BIN_THRESH).astype(float)\n    binarized = denoise(binarized, n=N, m=M)\n    \n    labeled = label(binarized)\n\n    df_comp = pd.DataFrame(regionprops_table(labeled, properties=[\"area\", \"bbox\"]))\n    \n    min_size = min(df_comp['area'].max() / 2, MIN_SIZE)\n    min_size = max(df_comp['area'].max() / 5, min_size)\n\n    crops, dists, to_ignore = component_distance(\n        viz, labeled, df_comp, plot=False, min_size=min_size\n    )\n\n    # Duplicates\n    duplicates = get_duplicates(dists, threshold=THRESHOLD_DUP)\n\n    get_clusters(df_comp, duplicates)\n\n    # Duplicates crop/rot\n    found_match_crop = False\n    if not len(duplicates) and len(df_comp) > 1:\n        df_comp['-area'] = - df_comp['area']\n        big_crops = df_comp.sort_values(\"-area\").index[:2]\n\n        x0, y0, x1, y1 = df_comp[[\"bbox-0\", \"bbox-1\", \"bbox-2\", \"bbox-3\"]].values[big_crops[0]]\n        mask = (labeled[x0:x1, y0:y1][..., None] == big_crops[0] + 1)\n        crop1 = (crops[big_crops[0]] * mask + 255 * (1 - mask)).astype(np.uint8)\n\n        x0, y0, x1, y1 = df_comp[[\"bbox-0\", \"bbox-1\", \"bbox-2\", \"bbox-3\"]].values[big_crops[1]]\n        mask =  (labeled[x0:x1, y0:y1][..., None] == big_crops[1] + 1)\n        crop2 = (crops[big_crops[1]] * mask + 255 * (1 - mask)).astype(np.uint8)\n        \n        found_match_crop = match_crops(crop1, crop2, threshold=THRESHOLD_DUP)\n        \n        if found_match_crop:\n            df_comp['cluster'][big_crops[0]] = 1\n            df_comp['cluster'][big_crops[1]] = 1\n\n            duplicates.append((big_crops[0], big_crops[1]))\n\n    # Crop based on matches\n    df_comp_kept = []\n    x0, x1, y0, y1 = None, None, None, None\n    if len(df_comp):\n        x0, x1, y0, y1, df_comp_kept = get_roi(\n            df_comp, viz.shape, found_crop=found_match_crop, min_size=min_size, margin=MARGIN\n        )\n        \n    # Refine with components dist\n    if (5 > len(df_comp_kept) > 1):\n\n        df_comp_kept = df_comp_kept.sort_values('area', ascending=False)\n        df_comp_kept = df_comp_kept[df_comp_kept['area'] > (df_comp_kept['area'].max() / 5)]\n        df_comp_kept['xc'] = (df_comp_kept['bbox-0'] + df_comp_kept['bbox-2']) / 2\n        df_comp_kept['yc'] = (df_comp_kept['bbox-1'] + df_comp_kept['bbox-3']) / 2\n        df_comp_kept['h'] = (df_comp_kept['bbox-2'] - df_comp_kept['bbox-0']) / 2\n        df_comp_kept['w'] = (df_comp_kept['bbox-3'] - df_comp_kept['bbox-1']) / 2\n        \n        dist_matrix = compute_dist_matrix(df_comp_kept.reset_index())\n    \n        dists = dist_matrix.min(0)\n        has_clusts = (dists.max() > 6)\n        \n        if has_clusts:\n            x0, y0, x1, y1 = df_comp_kept[['bbox-0', 'bbox-1', 'bbox-2', 'bbox-3']].values[0]\n            \n            if (viz[x0:x1, y0: y1].mean(-1) < 10).mean() > 0.05:  # black !\n                x0, y0, x1, y1 = df_comp_kept[['bbox-0', 'bbox-1', 'bbox-2', 'bbox-3']].values[1]\n        \n            shape = viz.shape\n            x0 = max(0, x0 - MARGIN)\n            y0 = max(0, y0 - MARGIN)\n            x1 = min(shape[0], x1 + MARGIN)\n            y1 = min(shape[1], y1 + MARGIN)\n    \n    if x0 is not None:  # Found a crop !\n        x0_ = int(x0 * factor)\n        x1_ = int(x1 * factor)\n        y0_ = int(y0 * factor)\n        y1_ = int(y1 * factor)\n\n        crop = image[x0_: x1_, y0_: y1_]\n\n    else:\n        crop = image\n    \n    del image\n    gc.collect()\n\n    # Final image\n    final_size = (np.array(crop.shape[:2]) / np.min(crop.shape[:2]) * IMG_SIZE).astype(int)[::-1]\n\n    img = cv2.resize(crop, final_size)\n    img = normalize(img, ratio)\n\n#     img = remove_bubble(img)\n#     img = remove_background(img.mean(-1).astype(int), img)\n\n    if PLOT:        \n        plt.figure(figsize=(10, 10))\n        plt.imshow(img)\n        plt.axis(False)\n        plt.title(image_id + \"  -  Deduped\")\n        plt.show()\n\n    path_d = SAVE_FOLDER_D + f\"{image_id}.png\"\n    if SAVE:\n        cv2.imwrite(path_d, img)\n\n    meta = test_df.iloc[idx].to_dict()\n    meta.update({\n        \"path\": path,\n        \"path_d\": path_d,\n    })\n    metadata.append(meta)\n\n    del img, crop\n    gc.collect()\n#     break\n\ndf = pd.DataFrame(metadata)\nif SAVE:\n    df.to_csv(\"df.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-10-04T08:01:53.153677Z","iopub.execute_input":"2022-10-04T08:01:53.154017Z","iopub.status.idle":"2022-10-04T08:04:03.18452Z","shell.execute_reply.started":"2022-10-04T08:01:53.153983Z","shell.execute_reply":"2022-10-04T08:04:03.183348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"markdown","source":"#### Params","metadata":{}},{"cell_type":"code","source":"EXP_FOLDERS = [  # auc 0.683 / loss 0.642\n    \"../input/mayo-weights-1/6/\",    # effnet-b0 smooth    - auc 0.660\n    \"../input/mayo-weights-1/16/\",   # resnet10 ranger     - auc 0.661\n    \"../input/mayo-weights-1/11/\",   # effnet-b0 w/batch   - auc 0.670\n]\n\nIMG_FOLDERS = [\n    SAVE_FOLDER,\n    SAVE_FOLDER_D,\n    SAVE_FOLDER_D,\n]\n\n\nUSE_TTA = True","metadata":{"execution":{"iopub.status.busy":"2022-10-04T08:04:03.186469Z","iopub.execute_input":"2022-10-04T08:04:03.186858Z","iopub.status.idle":"2022-10-04T08:04:03.194199Z","shell.execute_reply.started":"2022-10-04T08:04:03.18682Z","shell.execute_reply":"2022-10-04T08:04:03.192999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Main","metadata":{}},{"cell_type":"code","source":"predict_fct = predict_tta if USE_TTA else predict\nseed_everything(1337)\npreds_test = []\n\nfor exp_folder, img_folder in zip(EXP_FOLDERS, IMG_FOLDERS):\n    print(f'\\n- Exp {exp_folder}\\n')\n    \n    df = pd.read_csv(\"df.csv\", dtype={\"patient_id\": \"str\"})\n\n    if \"_d_\" in img_folder:\n        df[\"path\"] = df[\"path_d\"]\n    \n    df['target'] = 0\n    df[\"img_prop\"] = 1\n\n    config = Config(json.load(open(exp_folder + \"config.json\", \"r\")))\n\n    dataset = TileDataset(\n        df,\n        n_tiles=config.n_tiles_val,\n        transforms=get_transfos(augment=False, resize=config.resize),\n        train=False\n    )\n\n    model = define_model(\n        config.name,\n        num_classes=config.num_classes,\n        pretrained=False,\n        average=config.average,\n    ).cuda()\n\n    model.zero_grad()\n    model.eval()\n    \n    for w in sorted(glob.glob(exp_folder + \"*.pt\")):\n        load_model_weights(model, w, verbose=1)\n        \n        pred = predict_fct(model, dataset, config.loss_config, batch_size=16)\n        preds_test.append(pred)\n        \n#         print(pred)\n        \n    del model, dataset\n    gc.collect()\n    torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-10-04T08:08:46.169021Z","iopub.execute_input":"2022-10-04T08:08:46.169388Z","iopub.status.idle":"2022-10-04T08:08:58.9315Z","shell.execute_reply.started":"2022-10-04T08:08:46.169357Z","shell.execute_reply":"2022-10-04T08:08:58.930403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Post-process","metadata":{}},{"cell_type":"code","source":"preds = np.mean(preds_test, 0)\n\npreds = scale(preds, 0.15, 0.85)\npreds = np.clip(preds, 0.25, 0.75)\n\n# preds = np.where(np.abs(preds - 0.5) < 0.1, 1 - preds, preds)\npreds = 1 - preds","metadata":{"execution":{"iopub.status.busy":"2022-10-04T08:11:05.957824Z","iopub.execute_input":"2022-10-04T08:11:05.958757Z","iopub.status.idle":"2022-10-04T08:11:05.965391Z","shell.execute_reply.started":"2022-10-04T08:11:05.958706Z","shell.execute_reply":"2022-10-04T08:11:05.96443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Submission","metadata":{}},{"cell_type":"code","source":"df['pred'] = preds\n\ndf = df[[\"patient_id\", \"pred\"]].groupby(\"patient_id\").mean().reset_index()\n\nsub = pd.read_csv(DATA_PATH + \"sample_submission.csv\")[[\"patient_id\"]]\nsub = sub.merge(df)\n\nsub.columns = [\"patient_id\", \"LAA\"]\nsub[\"CE\"] = 1 - sub[\"LAA\"]\nsub = sub[['patient_id', 'CE', \"LAA\"]]\n\nsub.to_csv('submission.csv', index=False)\n\nsub.head()","metadata":{"execution":{"iopub.status.busy":"2022-10-04T08:19:35.989655Z","iopub.execute_input":"2022-10-04T08:19:35.990027Z","iopub.status.idle":"2022-10-04T08:19:36.016933Z","shell.execute_reply.started":"2022-10-04T08:19:35.989993Z","shell.execute_reply":"2022-10-04T08:19:36.015882Z"},"trusted":true},"execution_count":null,"outputs":[]}]}