{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":6774553,"sourceType":"datasetVersion","datasetId":3898019},{"sourceId":7070715,"sourceType":"datasetVersion","datasetId":4071818},{"sourceId":7071043,"sourceType":"datasetVersion","datasetId":4072075},{"sourceId":7092658,"sourceType":"datasetVersion","datasetId":4087402},{"sourceId":7329128,"sourceType":"datasetVersion","datasetId":4070763},{"sourceId":7329134,"sourceType":"datasetVersion","datasetId":4070766},{"sourceId":153398869,"sourceType":"kernelVersion"}],"dockerImageVersionId":30627,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# install the deb packages\n!yes | dpkg -i --force-depends /kaggle/input/pyvips-python-and-deb-package-gpu/linux_packages/archives/*.deb\n# install the python wrapper\n!pip install pyvips -f /kaggle/input/pyvips-python-and-deb-package-gpu/python_packages/ --no-index\n!pip install --no-index --no-deps ../input/xformers-wheel/xformers/*.whl","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-01-01T02:48:22.67579Z","iopub.execute_input":"2024-01-01T02:48:22.676107Z","iopub.status.idle":"2024-01-01T02:51:10.55438Z","shell.execute_reply.started":"2024-01-01T02:48:22.67608Z","shell.execute_reply":"2024-01-01T02:51:10.553217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pyvips\nfrom PIL import Image\nimport numpy as np\nfrom matplotlib import pyplot as plt\nimport cv2 as cv\nimport time\nimport torchstain as ts\nimport os\nimport gc\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset\nimport pandas as pd\nimport sys\nimport math\nfrom operator import itemgetter\nfrom numpy.lib.stride_tricks import sliding_window_view\nfrom itertools import compress\nimport h5py\nfrom torch.nn.functional import softmax\nfrom torch.utils.data import DataLoader\nfrom torch.nn import functional as F\nimport timm\nfrom xformers.components.feedforward import MLP\nfrom xformers.ops import fmha\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport faiss\nfrom copy import deepcopy\nfrom scipy import ndimage\nImage.MAX_IMAGE_PIXELS = None\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"max_split_size_mb:256\"\nos.environ['VIPS_DISC_THRESHOLD'] = '15gb' \ngc.enable()","metadata":{"execution":{"iopub.status.busy":"2024-01-01T02:56:01.957015Z","iopub.execute_input":"2024-01-01T02:56:01.957414Z","iopub.status.idle":"2024-01-01T02:56:01.967232Z","shell.execute_reply.started":"2024-01-01T02:56:01.957376Z","shell.execute_reply":"2024-01-01T02:56:01.966044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Preprocessing","metadata":{}},{"cell_type":"code","source":"def remove_duplicates_from_thumbnail(thumb_path, image_id):\n    thumbnail = os.path.join(thumb_path, str(image_id) + '_thumbnail.png')\n    tn = cv.imread(thumbnail)\n    tn = cv.cvtColor(tn, cv.COLOR_BGR2RGB)\n    # To remove the \"serrated\" border around the image\n    #kernel = cv.getStructuringElement(cv.MORPH_ELLIPSE, (3,3))\n    #tn = cv.erode(tn, kernel, iterations=1)\n    tn_h, tn_w = tn.shape[:2]\n    gray = cv.cvtColor(tn, cv.COLOR_BGR2GRAY)\n    _, thresh = cv.threshold(gray, 127, 255, 0)\n    contours, _ = cv.findContours(thresh, cv.RETR_EXTERNAL, cv.CHAIN_APPROX_SIMPLE)\n    blobs = []\n    coords = []\n    for cnt in contours:\n        x, y, w, h = cv.boundingRect(cnt)\n        if w > 100 and h > 100:\n            blobs.append(cnt)\n            coords.append((x, y, w, h))\n\n    similarities = []\n    for i in range(len(blobs)):\n        for j in range(i+1, len(blobs)):\n            sim = cv.matchShapes(blobs[i], blobs[j], 1, 0.0)\n            similarities.append((i, j, sim))\n\n    are_duplicates = all([x[2] < 0.2 for x in similarities]) if len(similarities) > 0 else False\n    if are_duplicates:\n        #max_sat = 0.0\n        max_pixels = 0\n        for i, coord in enumerate(coords):\n            print(i, coord)\n            x, y, w, h = coord\n            rect = tn[y:y+h, x:x+w]\n            #rect_hsv = cv.cvtColor(rect, cv.COLOR_BGR2HSV)\n            #mean_sat = np.mean(rect_hsv[:,:,1])\n            area = h * w\n            #print('Thumb', image_id, x, y, w, h, mean_sat)\n            #if (mean_sat > max_sat) and (mean_sat < 75):\n            if area > max_pixels:\n                #max_sat = mean_sat\n                max_pixels = area\n                max_sat_rect = rect\n                max_sat_idx = i\n        return max_sat_rect, coords[max_sat_idx], tn_w, tn_h\n    else:\n        return tn, [0, 0, 0, 0], tn_w, tn_h\n\ndef resize_thumbnail(tn, wsi_w, tn_w):\n    eps = 0.0002 # needed for edge cases that have one pixel missing upon resizing\n    scale = wsi_w / tn_w\n    closest_scale_2 = [2**x for x in range(6) if 2**x < scale][-1]\n    tn = cv.resize(tn, None, fx= scale/closest_scale_2 + eps, fy= scale/closest_scale_2 + eps)\n    return tn, closest_scale_2\n\ndef select_stain_normalizer():\n    ref_path = '/kaggle/input/stain-targets'\n    #ref_patches = [13568, 17637, 36583, 50932]\n    #ref_patches = [17637, 31594, 48734, 50932]\n    ref_patches = [31594, 36583, 40864, 50932]\n    ref = np.random.choice(len(ref_patches), 1)[0]\n    stain_normalizer = ts.normalizers.MacenkoNormalizer(backend='torch') #if np.random.random() < 0.5 else ts.normalizers.ReinhardNormalizer(backend='torch')\n    target = cv.imread(os.path.join(ref_path, str(ref_patches[ref]) + '_ref.png'), cv.COLOR_RGBA2RGB)\n    target = np.moveaxis(target, -1, 0)\n    target = torch.from_numpy(target)\n    stain_normalizer.fit(target)\n    return stain_normalizer\n\ndef countBlackPixels(patch):\n    black = cv.inRange(patch, (0, 0, 0), (20, 20, 20))\n    return cv.countNonZero(black)\n\ndef countWhitePixels(patch):\n    white = cv.inRange(patch, (210, 210, 210), (255, 255, 255))\n    return cv.countNonZero(white)\n\ndef countGrayPixels(patch):\n    gray_patch = cv.cvtColor(patch, cv.COLOR_RGB2HSV)\n    gray = cv.inRange(gray_patch, (0, 0, 127), (179, 10, 255))\n    return cv.countNonZero(gray)\n\ndef isWhiteBlackPatch(patch, threshold):\n    w, h = patch.shape[:2]\n    num_pixels = w*h\n    num_white = countWhitePixels(patch)\n    num_black = countBlackPixels(patch)\n    num_gray = countGrayPixels(patch)\n    return True if ((max(num_white, num_gray) + num_black) / num_pixels) > threshold else False\n\ndef isPurplishGrayPatch(patch):\n    # This is based on TMA 91.png; purplish gray\n    w, h = patch.shape[:2]\n    num_pixels = w*h\n    gray = cv.inRange(patch, (195, 185, 205), (215, 205, 225))\n    num_gray = cv.countNonZero(gray)\n    #print(num_gray, num_pixels, num_gray / num_pixels)\n    return True if num_gray / num_pixels > 0.5 else False\n\ndef isPaleEosinPatch(patch):\n    # Based on trial and error on observation from slides\n    std_r = np.std(patch[:, :, 0])\n    std_g = np.std(patch[:, :, 1])\n    std_b = np.std(patch[:, :, 2])\n    #print('Var', std_r, std_g, std_b)\n    return (std_r < 10 and std_g < 15 and std_b < 10)\n\ndef proportionHematoxylin(patch):\n    patch = patch.astype(np.uint8)\n    w, h = patch.shape[:2]\n    num_pixels = w*h\n    img = cv.cvtColor(patch, cv.COLOR_RGB2HSV)\n    is_purple = (img[:, :, 0] > 128) | (img[:, :, 0] < 132) # This hue range correspond mostly with hematoxylin\n    mask = is_purple & (img[:, :, 1] > 50) & (img[:, :, 2] > 200) # We want to keep the highest saturated purple pixels\n    return np.sum(mask) / num_pixels\n    #histr = cv.calcHist([img],[0],None,[256],[0,256])\n    #num_hx = histr[130] + histr[131] # \n    #return num_hx / num_pixels\n\ndef proportionEosin(patch):\n    patch = patch.astype(np.uint8)\n    w, h = patch.shape[:2]\n    num_pixels = w*h\n    \n    return np.sum(np.all(patch==255, axis=2))  / num_pixels\n\ndef padTile(tn, top, bottom, left, right, height, width):\n    print('Surprise!', top, bottom, left, right, height, width)\n    pad_t = pad_b = pad_l = pad_r = 0\n    if top < 0:\n        pad_t = -top\n        top = 0\n    if bottom > height:\n        pad_b = bottom - height\n        bottom = height\n    if left < 0:\n        pad_l = -left\n        left = 0\n    if right > width:\n        pad_r = right - width\n        right = width\n    part_image = tn[top:bottom, left:right] if isinstance(tn, np.ndarray) else tn.crop(left, top, right - left, bottom - top).numpy()\n    padded = cv.copyMakeBorder(part_image, pad_t, pad_b, pad_l, pad_r, cv.BORDER_CONSTANT, None, [0, 0, 0])\n    return padded\n\ndef clusterTiles(selected, height, width, bag):\n    tile_scores = {}\n    for tile in selected:\n        tile_h, tile_w = tile\n        tile_scores[tile] = tile_scores.get(tile, 0) + 3\n        offset = np.meshgrid(np.arange(-1, 2), np.arange(-1, 2))\n        neighbors = (np.array(offset).T.reshape(-1, 2) + [tile_h, tile_w]).tolist()\n        neighbors.pop(neighbors.index([tile_h, tile_w]))\n        for neighbor in neighbors:\n            if (neighbor[0] > -1) & (neighbor[1] > -1) & (neighbor[0] < height) & (neighbor[1] < width):\n                tile_scores[tuple(neighbor)] = tile_scores.get(tuple(neighbor), 0) + 2\n    clustered = [x for x, y in tile_scores.items() if y > 3]\n    borderline = [x for x, y in tile_scores.items() if y == 3]\n    for tile in borderline:\n        tile_h, tile_w = tile\n        offset = np.meshgrid(np.arange(-1, 2), np.arange(-1, 2))\n        neighbors = (np.array(offset).T.reshape(-1, 2) + [tile_h, tile_w]).tolist()\n        if [x for x in neighbors if tuple(x) in clustered]:\n            clustered.append(tile)\n    clustered = sorted(clustered)\n    print('Chosen', len(selected), len(clustered))\n    return clustered\n\ndef mapTiles(image, width, height, w_tiles, h_tiles, side, tiles_bag, stain_normalizer):\n    #w_tiles = width // side\n    #h_tiles = height // side\n    #print('Stain', h_tiles, w_tiles)\n    grid_info = {'init_w_tiles': w_tiles,\n                 'init_h_tiles': h_tiles,\n                 'shrunk_w_tiles': 0,\n                 'shrunk_h_tiles': 0,\n                 'pad_w': 0,\n                 'pad_h': 0}\n    border_side = (width - w_tiles * side) // 2\n    border_height = (height - h_tiles * side) // 2\n    shrink = np.sqrt(15000/(h_tiles*w_tiles))\n    if shrink < 1.0:\n        tiles_bag = int(np.floor(tiles_bag / shrink))\n    print(width, height, w_tiles, h_tiles, border_side, border_height, side, tiles_bag, shrink)\n    patches = []\n    coords = []\n    h_values = []\n    e_values = []\n    minH = 1000000\n\n    for h in range(h_tiles):\n        for w in range(w_tiles):\n            mask = False\n            top = border_height + h * side\n            bottom = border_height + (h + 1) * side\n            left = border_side + w * side\n            right = border_side + (w + 1) * side\n            if any([top < 0, bottom > height, left < 0, right > width]):\n                patch_rgb = padTile(image, top, bottom, left, right, height, width)\n            else:\n                patch_rgb = image[top:bottom, left:right]\n            #patch_rgb = image[(border_height + h * side):(border_height + (h + 1) * side), (border_side + w * side):(border_side + (w + 1) * side)]\n            patchDev = np.sum(np.std(patch_rgb, axis=(0,1)))\n            isWhiteBlack = False if side == 704 else isWhiteBlackPatch(patch_rgb, 0.20)\n            mask = patchDev < 20.0 or isWhiteBlack or isPurplishGrayPatch(patch_rgb) #or isPaleEosinPatch(patch_rgb)\n            if not mask:\n                patch_rgb = np.moveaxis(patch_rgb, -1, 0)\n                patch_rgb = torch.from_numpy(patch_rgb)\n                try:\n                    patch_rgb, H, E = stain_normalizer.normalize(I=patch_rgb, stains=True)\n                except (torch.linalg.LinAlgError, IndexError):\n                    continue\n                patch_rgb = patch_rgb.numpy().astype(np.uint8)\n                propH = proportionHematoxylin(H.numpy().astype(np.uint8))\n                propE = proportionEosin(E.numpy().astype(np.uint8))\n                if side == 704:\n                    patch_rgb = cv.resize(patch_rgb, (352, 352))\n                e_values.append(propE)\n                h_values.append(propH)\n                patches.append(patch_rgb)\n                coords.append((h, w))\n    \n    e_list = list(zip(coords, e_values))\n    h_list = list(zip(coords, h_values))\n    e_list = sorted(e_list, key=lambda x: x[1], reverse=True)\n    h_list = sorted(h_list, key=lambda x: x[1], reverse=True)\n    e_rank = [(k, idx) for idx, (k, v) in enumerate(e_list)]\n    h_rank = [(k, idx) for idx, (k, v) in enumerate(h_list)]\n    e_rank = sorted(e_rank, key=lambda x: (x[0], x[1]))\n    h_rank = sorted(h_rank, key=lambda x: (x[0], x[1]))\n    ranking = [(x[0], (x[1] + y[1]) / 2) for x, y in list(zip(e_rank, h_rank))]\n    ranking = sorted(ranking, key=lambda x: x[1])\n    coords = [x[0] for x in ranking[:tiles_bag]]\n    coords = sorted(coords, key=lambda x: (x[0], x[1]))\n    if side != 704:\n        coords = clusterTiles(coords, h_tiles, w_tiles, tiles_bag)\n        pop_idx = []\n        for i, tile in enumerate(coords):\n            patch_check = image[(border_height + tile[0] * side):(border_height + (tile[0] + 1) * side), \n                                (border_side + tile[1] * side):(border_side + (tile[1] + 1) * side)]\n            patchDev = np.sum(np.std(patch_check, axis=(0,1)))\n            if any([patchDev < 20.0, isWhiteBlackPatch(patch_check, 0.40), isPurplishGrayPatch(patch_check)]):\n                pop_idx.append(i)\n        coords = [x for i, x in enumerate(coords) if i not in pop_idx]\n    \n    if shrink < 0.98: # To avoid situations where bincount < init_w/h_tiles\n        new_height = int(np.floor(h_tiles * shrink))\n        new_width = int(np.floor(w_tiles * shrink))\n        h_tiles = new_height if new_height % 2 == 0 else max(new_height - 1, 4)\n        w_tiles = new_width if (new_width % 4 == 0) else max(new_width - (new_width % 4), 4)\n        width_count = np.bincount([x[1] for x in coords])\n        height_count = np.bincount([x[0] for x in coords])\n        first_w = np.argmax(sliding_window_view(width_count, np.min([len(width_count), w_tiles])).sum(axis=-1))\n        first_h = np.argmax(sliding_window_view(height_count, np.min([len(height_count), h_tiles])).sum(axis=-1))\n        grid_info['shrunk_w_tiles'] = w_tiles\n        grid_info['shrunk_h_tiles'] = h_tiles\n        grid_info['pad_w'] = first_w\n        grid_info['pad_h'] = first_h\n        coords = [(x - first_h, y - first_w) for x, y in coords if y in list(range(first_w, np.min([len(width_count), w_tiles]) + first_w)) \n                  and x in list(range(first_h, np.min([len(height_count), h_tiles]) + first_h))]\n    \n    if len(coords) > 450:\n        margin = len(coords) // 2\n        coords = coords[margin:(margin + 450)]\n    \n    coords.append((h_tiles, w_tiles))\n    return patches, coords, grid_info\n\ndef setNumberTiles(width, height, side):\n    # Need to ensure that the number of tiles is divisible by 8 (required for attn_bias from xFormers)\n    # Also, minimum side size set to 4 for TMAs\n    width_tiles = width // side\n    height_tiles = height // side\n    width_tiles = width_tiles if (width_tiles % 2 == 0) else max(width_tiles - 1, 4)\n    height_tiles = height_tiles if (height_tiles % 4 == 0) else max(height_tiles - (height_tiles % 4), 4)\n    return width_tiles, height_tiles\n\ndef extractTiles(image, scaled_coord, factor, side, coords, grid_info, stain_normalizer):\n    wsi_height, wsi_width = image.height, image.width\n    if np.sum(scaled_coord) == 0:\n        border_side = (wsi_width - grid_info['init_w_tiles'] * side) // 2 + grid_info['pad_w'] * side\n        border_height = (wsi_height - grid_info['init_h_tiles'] * side) // 2 + grid_info['pad_h'] * side\n    else:\n        x, y, width, height = scaled_coord\n        roi_width = round(width * factor)\n        roi_height = round(height * factor)\n        border_side = round(x * factor)\n        border_height = round(y * factor)\n    #print(grid_info)\n    #border_side = round(x * factor) + (image.width - grid_info['init_w_tiles'] * side) // 2 + grid_info['pad_w'] * side\n    #border_height = round(y * factor) + (image.height - grid_info['init_h_tiles'] * side) // 2 + grid_info['pad_h'] * side\n    #print(roi_width, roi_height, w_tiles, h_tiles, border_side, border_height, side)\n    patches = []\n    copy_coords = deepcopy(coords)\n    for h, w in copy_coords[:-1]:\n        #patch = np.array(image.read_region((border_side + w * side, border_height + h * side), 0, size=(side, side)))top = border_height + h * side\n        top = border_height + h * side\n        bottom = border_height + (h + 1) * side\n        left = border_side + w * side\n        right = border_side + (w + 1) * side\n        if any([top < 0, bottom > wsi_height, left < 0, right > wsi_width]):\n            patch = padTile(image, top, bottom, left, right, wsi_height, wsi_width)\n        else:\n            patch = image.crop(left, top, side, side).numpy()\n        #patch = image.crop(border_side + w * side, border_height + h * side, side, side).numpy()\n        patch_rgb = cv.cvtColor(patch, cv.COLOR_RGBA2RGB)\n        #print('White', h, w, countWhitePixels(patch_rgb))\n        patch_rgb = np.moveaxis(patch_rgb, -1, 0)\n        patch_rgb = torch.from_numpy(patch_rgb)\n        try:\n            patch_rgb, _, _ = stain_normalizer.normalize(I=patch_rgb, stains=True)\n            patch_rgb = patch_rgb.numpy().astype(np.uint8)\n            patches.append(patch_rgb)\n        except:\n            print(h, w)\n            print('Coords', border_height + h * side, border_side + w * side, border_height, border_side, h, w)\n            coords.pop(coords.index((h, w)))\n            #plt.imshow(np.moveaxis(patch_rgb.numpy(), 0, -1))\n            #sys.exit(0)\n    del image\n    gc.collect()\n    return patches, coords\n\ndef createTiles(slide_path, thumb_path, image_id, tile_side):\n    stain_normalizer = select_stain_normalizer()\n    wsi = pyvips.Image.new_from_file(os.path.join(slide_path, str(image_id) + '.png'))\n    wsi_w = wsi.width\n    wsi_h = wsi.height\n    if tile_side == 352:\n        tiles_bag = 192\n        roi, scaled_coord, tn_w, tn_h = remove_duplicates_from_thumbnail(thumb_path, image_id)\n        roi, tn_factor = resize_thumbnail(roi, wsi_w, tn_w)\n        roi_h, roi_w = roi.shape[:2]\n        roi_side = tile_side // tn_factor\n        if np.sum(scaled_coord) == 0:\n            w_tiles, h_tiles = setNumberTiles(wsi_w, wsi_h, tile_side)\n        else:\n            scaled_w = round(roi_w * tn_factor)\n            scaled_h = round(roi_h * tn_factor)\n            w_tiles, h_tiles = setNumberTiles(scaled_w, scaled_h, tile_side)\n        _, coords, grid_info = mapTiles(roi, roi_w, roi_h, w_tiles, h_tiles, roi_side, tiles_bag, stain_normalizer)\n        tiles, coords = extractTiles(wsi, scaled_coord, wsi_w/tn_w, tile_side, coords, grid_info, stain_normalizer)\n    else:  \n        tiles_bag = 48\n        #wsi = open_slide(os.path.join(slide_path, str(image_id) + '.png'))\n        #wsi = np.array(wsi.read_region((0, 0), 0, size=(wsi_w, wsi_h)))\n        wsi = wsi.crop(0, 0, wsi_w, wsi_h).numpy()\n        wsi = cv.cvtColor(wsi, cv.COLOR_RGBA2RGB)\n        w_tiles, h_tiles = setNumberTiles(wsi_w, wsi_h, tile_side)\n        tiles, coords, _ = mapTiles(wsi, wsi_w, wsi_h, w_tiles, h_tiles, tile_side, tiles_bag, stain_normalizer)\n    #print('Done', len(tiles), len(coords))\n    del wsi\n    gc.collect()\n    return tiles, coords","metadata":{"execution":{"iopub.status.busy":"2024-01-01T02:51:19.46432Z","iopub.execute_input":"2024-01-01T02:51:19.465033Z","iopub.status.idle":"2024-01-01T02:51:19.54479Z","shell.execute_reply.started":"2024-01-01T02:51:19.464972Z","shell.execute_reply":"2024-01-01T02:51:19.543318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Dataset + utils","metadata":{}},{"cell_type":"code","source":"class HistoDataset(Dataset):\n    def __init__(self, img_path, thumb_path, sample_list, img_transforms = None):\n        self.img_path = img_path\n        self.thumb_path = thumb_path\n        self.sample_list = sample_list\n        self.metadata = pd.read_csv('/kaggle/input/UBC-OCEAN/test.csv')\n        #self.metadata = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\n        self.img_transforms = img_transforms\n    \n    def __len__(self):\n        return len(self.sample_list)\n\n    def __getitem__(self, index):\n        image_id = self.sample_list[index]\n        print('Preparing tiles for ', image_id)\n        wsi_w = self.metadata.loc[self.metadata['image_id']==image_id, 'image_width'].values[0]\n        wsi_h = self.metadata.loc[self.metadata['image_id']==image_id, 'image_height'].values[0]\n        tile_side = 704 if wsi_w * wsi_h < 30000000 else 352\n        tiles, coords = createTiles(self.img_path, self.thumb_path, image_id, tile_side)\n        if tile_side == 352:\n            try:\n                max_h = max(coords, key=lambda x: x[0])[0]\n                max_w = max(coords, key=lambda x: x[1])[1]\n                coords_map = np.zeros((max_h, max_w))\n                for h, w in coords[:-1]:\n                    coords_map[h, w] = 1\n                struct = [[1, 1, 1], [1, 1, 1], [1, 1, 1]]\n                clustered, _ = ndimage.label(coords_map, struct)\n                _, c_size = np.unique(clustered, return_counts=True)\n                target_clusters = [(idx, size) for idx, size in enumerate(c_size) if (size > 11) and (size < 51)]\n                big_clusters = [(idx + 1, size) for idx, size in enumerate(c_size[1:]) if (size > 50) and (size < 129)]\n                small_clusters = [(idx, size) for idx, size in enumerate(c_size) if (size < 12) and (size > 7)]\n                if len(target_clusters) > 0:\n                    select_cluster = max(target_clusters, key=lambda x: x[1])[0]\n                elif len(big_clusters) > 0:\n                    select_cluster = min(big_clusters, key=lambda x: x[1])[0]\n                else:\n                    select_cluster = max(small_clusters, key=lambda x: x[1])[0]\n                select_tiles = np.where(clustered == select_cluster)\n                select_tiles = list(tuple(zip(*select_tiles)))\n                #print('Tiles coords', len(tiles), len(coords))\n                select_coords = [idx for idx, (x, y) in enumerate(coords[:-1]) if (x, y) in select_tiles]\n                #print('Select', select_coords)\n                new_tiles = [tiles[i] for i in select_coords]\n                min_h = min(select_tiles, key=lambda x: x[0])[0]\n                min_w = min(select_tiles, key=lambda x: x[1])[1]\n                select_tiles = [(x - min_h, y - min_w) for x, y in select_tiles]\n                new_h_tiles = max(select_tiles, key=lambda x: x[0])[0] + 1\n                new_w_tiles = max(select_tiles, key=lambda x: x[1])[1] + 1\n                new_h_tiles = new_h_tiles if new_h_tiles % 2 == 0 else new_h_tiles + 1\n                new_w_tiles = new_w_tiles if (new_w_tiles % 4 == 0) else new_w_tiles + (4 - new_w_tiles % 4)\n                select_tiles = select_tiles + [(new_h_tiles, new_w_tiles)]\n                tiles = new_tiles\n                coords = select_tiles\n            except:\n                tiles = torch.rand(1, 3, 352, 352)\n                coords = [(0, 0), (50, 50)]\n        if self.img_transforms is not None:\n            trans_tiles = []\n            drop_coords = []\n            try:\n                for i, tile in enumerate(tiles):\n                    tile = self.img_transforms(image=tile)['image']\n                    trans_tiles.append(tile)\n                tiles = torch.stack(trans_tiles)\n            except:\n                tiles = torch.rand(1, 3, 352, 352)\n                coords = [(0, 0), (50, 50)]\n        return (tiles, torch.tensor(coords), image_id)\n\ndef get_loaders(\n        img_path,\n        thumb_path,\n        query_sample_list,\n        batch_size,\n        val_transform,\n        num_workers=0,\n        pin_memory=False\n):\n    query_dataset = HistoDataset(\n        img_path = img_path,\n        thumb_path = thumb_path,\n        sample_list = query_sample_list,\n        img_transforms = val_transform\n    )\n\n    query_loader = DataLoader(\n        query_dataset,\n        batch_size=batch_size,\n        num_workers=num_workers,\n        pin_memory=pin_memory,\n        shuffle=False\n    )\n\n    return query_loader\n\ndef load_checkpoint(checkpoint, model):\n    print('Loading checkpoint...')\n    model.load_state_dict(checkpoint['state_dict'])","metadata":{"execution":{"iopub.status.busy":"2024-01-01T02:56:24.89651Z","iopub.execute_input":"2024-01-01T02:56:24.897478Z","iopub.status.idle":"2024-01-01T02:56:24.920886Z","shell.execute_reply.started":"2024-01-01T02:56:24.897441Z","shell.execute_reply":"2024-01-01T02:56:24.919806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Models","metadata":{}},{"cell_type":"code","source":"def positional_encoding(dim, length, base):\n    assert dim % 2 == 0, f\"Cannot use sin/cos positional encoding with odd dimension (got dim={dim})\"\n\n    pe = torch.zeros(length, dim)\n    position = torch.arange(0, length, dtype=torch.float).unsqueeze(1)\n    div_term = torch.exp(torch.arange(0, dim, 2, dtype=torch.float) * -(math.log(base) / dim))\n    pe[:, 0::2] = torch.sin(position * div_term)\n    pe[:, 1::2] = torch.cos(position * div_term)\n\n    return pe.to(dtype=torch.float16)\n\ndef apply_rotary_position_embeddings(pos_enc, *tensors):\n    assert len(tensors) > 0, \"At least one input tensor is required\"\n    \n    cos_pos = pos_enc[..., 1::2].repeat_interleave(2, 1)\n    sin_pos = pos_enc[..., 0::2].repeat_interleave(2, 1)\n    \n    cos_pos = cos_pos.expand_as(tensors[0])\n    sin_pos = sin_pos.expand_as(tensors[0])\n\n    outputs = []\n    for t in tensors:\n        t_r = torch.empty_like(t)\n        t_r[..., 0::2] = -t[..., 1::2]\n        t_r[..., 1::2] = t[..., 0::2]\n        outputs.append(t * cos_pos + t_r * sin_pos)\n\n    return outputs if len(tensors) > 1 else outputs[0]\n\nclass Rotary2D:\n    def __init__(self, dim, base = 10000):\n        self.dim = dim\n        self.base = base\n        self.pos_cached = None\n        self.w_size_cached = None\n        self.h_size_cached = None\n\n    def __call__(self, x, grid_h, grid_w):\n        assert grid_h % 2 == 0, 'Grid height needs to be an even number'\n        assert grid_w % 2 == 0, 'Grid width needs to be an even number'\n        if self.pos_cached is None or self.w_size_cached != grid_w or self.h_size_cached != grid_h:\n            self.h_size_cached = grid_h\n            self.w_size_cached = grid_w\n\n            position_x = positional_encoding(grid_h, self.dim // 2, self.base)\n            position_y = positional_encoding(grid_w, self.dim // 2, self.base)\n\n            position_x = position_x.reshape(grid_h, -1, 2)\n            position_y = position_y.reshape(grid_w, -1, 2)\n\n            self.pos_cached = torch.empty(grid_h * grid_w, self.dim, dtype=torch.float16, device=x.device)\n            for i in range(grid_h):\n                for j in range(grid_w):\n                    emb = torch.cat([\n                        position_x[i, 0::2],\n                        position_y[j, 0::2],\n                        position_x[i, 1::2],\n                        position_y[j, 1::2]], 0).flatten(-2)\n                    self.pos_cached[i * grid_w + j] = emb.to(x.dtype).to(x.device)\n        return self.pos_cached\n\nclass RoFormerLayer(nn.Module):\n    def __init__(self, features, heads, dropout=0.0):\n        super(RoFormerLayer, self).__init__()\n        self.features = features\n        self.heads = heads\n        self.head_dim = self.features // heads\n        self.dropout = dropout\n\n        self.norm1 = nn.LayerNorm(self.features)\n        self.norm2 = nn.LayerNorm(self.features)\n        self.qkv_proj = nn.Linear(self.features, 3 * self.features)\n        self.rope = Rotary2D(features)\n        self.mlp = MLP(features, dropout=dropout, activation='relu', hidden_layer_multiplier=4)\n\n    def _add_mask(self, embed, coords, heads):\n        height, width = coords[-1]\n        #print('Coords', coords.shape, height, width)\n        batch_size, _, feat_size = embed.shape\n        #print('Embed', embed.shape)\n        new_embed = torch.zeros(batch_size, height * width, feat_size).to(dtype=embed.dtype).to(device=embed.device)\n        pre_mask = torch.full((batch_size, height*width), -math.inf).to(dtype=embed.dtype).to(device=embed.device)\n        pre_mask.requires_grad = False\n        for i, (x, y) in enumerate(coords[:-1]):\n            new_embed[:, x * width + y] = embed[:, i]\n        mask, _ = fmha.BlockDiagonalMask.from_tensor_list([pre_mask])\n        \n        return new_embed, mask\n\n    def forward(self, x, coords):\n        grid_h, grid_w = coords[-1]\n        x, mask = self._add_mask(x, coords, self.heads) # (B, H*W, D)\n        h = self.norm1(x) # (B, H*W, D)\n        bs, n = h.shape[:2]\n        qkv = self.qkv_proj(h) # (B, H*W, 3*D)\n        q, k, v = qkv.chunk(3, dim=-1) # (B, H*W, D)\n        q, k = apply_rotary_position_embeddings(self.rope(h, grid_h, grid_w), q, k)\n        q, k, v = q.reshape(bs, n, self.heads, self.head_dim), k.reshape(bs, n, self.heads, self.head_dim), v.reshape(bs, n, self.heads, self.head_dim)\n        att = fmha.memory_efficient_attention(q, k, v, attn_bias=mask, p = self.dropout, op=(fmha.cutlass.FwOp, fmha.cutlass.BwOp))\n        o = self.norm2(h + att.reshape(bs, n, h.size(-1)))\n        ff = self.mlp(o)\n        out = o + ff\n        return out, mask\n\nclass AttentionMIL(nn.Module):\n    def __init__(self, in_features, hidden_features, classes = 5, is_gated=False, is_inference=False):\n        super(AttentionMIL, self).__init__()\n        self.L = in_features\n        self.D = hidden_features\n        self.K = 1\n        self.classes = classes\n        self.is_gated = is_gated\n        self.is_inference = is_inference\n\n        self.attention_V = nn.Sequential(\n            nn.Linear(self.L, self.D),\n            nn.Tanh()\n        )\n\n        if self.is_gated:\n            self.attention_U = nn.Sequential(\n                nn.Linear(self.L, self.D),\n                nn.Sigmoid()\n            )\n        \n        self.attention_weights = nn.Linear(self.D, self.K)\n        self.classifier = nn.Sequential(\n            nn.Linear(self.L*self.K, self.classes)\n        )\n\n    def forward(self, x):\n        attention = self.attention_V(x)\n        if self.is_gated:\n            attention_u = self.attention_U(x)\n            attention = attention * attention_u\n        \n        attention = self.attention_weights(attention)\n\n        if self.is_inference:\n            pass\n        else:\n            attention = torch.transpose(attention, 2, 1)\n            attention = F.softmax(attention, dim = -1)\n            multiply_layer = torch.bmm(attention, x)\n            y_logit = self.classifier(multiply_layer)\n            y_hat = torch.argmax(F.softmax(y_logit, dim=-1), dim=-1)\n\n            return y_logit, y_hat, multiply_layer\n\nclass MultiHeadMIL(nn.Module):\n    def __init__(self, in_features, hidden_features, attn_features, classes = 5, heads = 4, dropout = 0.0, is_inference=False):\n        super(MultiHeadMIL, self).__init__()\n        self.features = in_features\n        self.hidden = hidden_features\n        self.attn = attn_features\n        self.classes = classes\n        self.heads = heads\n        self.is_inference = is_inference\n        \n        self.input_projection = nn.Sequential(\n            nn.Linear(self.features, self.hidden),\n            nn.ReLU(),\n            nn.Dropout(dropout)\n        )\n        self.attn_projection = nn.Linear(self.hidden, self.attn)\n        self.output_projection = nn.Linear(self.attn, self.hidden)\n\n        self.class_tokens = nn.Parameter(nn.init.xavier_normal_(torch.rand((1, classes, attn_features), requires_grad=True)))\n\n        self.classifier = nn.ModuleList(\n            [nn.Linear(self.hidden, 1) for _ in range(self.classes)]\n        )\n\n    def forward(self, x, mask):\n        bs, n = x.shape[:2]\n        proj = self.input_projection(x)\n\n        keys = self.attn_projection(proj)\n        values = self.attn_projection(proj)\n\n        k = keys.view(bs, n, self.heads, self.attn // self.heads)\n        v = values.view(bs, n, self.heads, self.attn // self.heads)\n        q = self.class_tokens.view(bs, self.classes, self.heads, self.attn // self.heads).repeat(1, len(mask._batch_sizes), 1, 1).to(k)\n        \n        class_mask = fmha.BlockDiagonalMask.from_seqlens(q_seqlen=[self.classes] * len(mask._batch_sizes),\n                                                         kv_seqlen= [i[1] - i[0] for i in mask.k_seqinfo.intervals()])\n\n        att = fmha.memory_efficient_attention(q, k, v, attn_bias=class_mask, op=(fmha.cutlass.FwOp, fmha.cutlass.BwOp))\n        att = att.view(bs, self.classes, self.attn)\n        out = self.output_projection(att)\n        if self.is_inference:\n            pass\n        else:\n            y_logit = torch.hstack([self.classifier[label](out[:, label]) for label in range(self.classes)])\n            y_hat = torch.argmax(F.softmax(y_logit, dim=-1), dim=-1)\n\n            return y_logit, y_hat, out\n\n\nclass WSINet(nn.Module):\n    def __init__(self, \n                 backbone='resnet50d.ra2_in1k', \n                 in_channels = 3,  \n                 num_classes = 5,\n                 heads = 4,\n                 dropout = 0.0,\n                 pretrained=True,\n                 pretrained_cfg_overlay=None,\n                 is_inference=False):\n        super(WSINet, self).__init__()\n\n        self.num_classes = num_classes\n        self.encoder = timm.create_model(\n                    backbone,\n                    in_chans = in_channels,\n                    features_only = False,\n                    pretrained = pretrained,\n                    pretrained_cfg_overlay = pretrained_cfg_overlay\n        )\n        features = self.encoder.get_classifier().in_features\n        self.encoder.reset_classifier(0)\n        for param in self.encoder.parameters():\n            param.requires_grad = False\n        self.roformer = RoFormerLayer(features, heads=heads, dropout=dropout)\n        self.attention = AttentionMIL(features, 512, self.num_classes, is_gated=True, is_inference=is_inference)\n\n    def forward(self, x, coords):\n        features = self.encoder(x).unsqueeze(0)\n        features, mask = self.roformer(features, coords)\n        y_prob, y_hat, attention = self.attention(features)\n        \n        return y_prob, y_hat, attention\n    \ndef compute_centroids(support_features, support_labels):\n    n_labels = len(torch.unique(support_labels))\n    \n    return torch.cat(\n        [support_features[torch.nonzero(support_labels == label)].mean(0) for label in range(n_labels)]\n    )\n\nclass OSLO():\n    def __init__(\n        self,\n        inference_steps,\n        lambda_s,\n        lambda_z,\n        ema_weight = 1.0,\n        use_inlier_latent = True\n    ):\n        super().__init__()\n        self.inference_steps = inference_steps\n        self.lambda_s = lambda_s\n        self.lambda_z = lambda_z\n        self.ema_weight = ema_weight\n        self.use_inlier_latent = use_inlier_latent\n\n    def cosine(self, X, Y):\n        return F.normalize(X, dim=-1) @ F.normalize(Y, dim=-1).T\n\n    def get_logits(self, centroids, query_features):\n        return self.cosine(query_features, centroids)  # [query_size, num_classes]\n#         logits = torch.empty((query_features.size(0), centroids.size(0)))\n#         for i in range(query_features.size(0)):\n#             for j in range(centroids.size(0)):\n#                 logits[i, j] = torch.sqrt(torch.sum((query_features[i] - centroids[j])**2))\n#         return logits\n\n    def __call__(\n        self,\n        support_features,\n        query_features,\n        support_labels,\n        **kwargs,\n    ):\n        num_classes = support_labels.unique().size(0)\n        one_hot_labels = F.one_hot(support_labels, num_classes)  # [support_size, num_classes]\n        print('OH', one_hot_labels.size())\n        support_size = support_features.size(0)\n\n        centroids = compute_centroids(support_features, support_labels)  # [num_classes, feature_dim]\n        print('Centroids', centroids)\n        latent_targets = (1 / num_classes) * torch.ones(query_features.size(0), num_classes)  # [query_size, num_classes]\n        inlier_scores = 0.5 * torch.ones((query_features.size(0), 1))\n        print('Features', query_features)\n\n        for _ in range(self.inference_steps):\n            # Compute inlier scores\n            logits_q = self.get_logits(centroids, query_features)  # [query_size, num_classes]\n            print('Logits q', logits_q)\n            inlier_scores = (\n                self.ema_weight * ((latent_targets * logits_q / self.lambda_s).sum(-1, keepdim=True).sigmoid())\n                + (1 - self.ema_weight) * inlier_scores\n            )  # [query_size, 1]\n            print('Inlier', inlier_scores)\n            # Compute new latent targets\n            latent_targets = (\n                (self.ema_weight * ((inlier_scores * logits_q / self.lambda_z).softmax(-1))\n                    + (1 - self.ema_weight) * latent_targets\n                )\n                if self.use_inlier_latent\n                else (\n                    self.ema_weight * ((logits_q / self.lambda_z).softmax(-1))\n                    + (1 - self.ema_weight) * latent_targets\n                )\n            )  # [query_size, num_classes]\n            print('Targets', latent_targets.size())\n            outlier_scores = 1.0 - inlier_scores\n\n            # Compute new centroids\n            all_features = torch.cat([support_features, query_features], 0)  # [support_size + query_size, feature_dim]\n            all_targets = torch.cat([one_hot_labels, latent_targets], dim=0)  # [support_size + query_size, num_classes]\n            all_inliers_scores = (\n                torch.cat([torch.ones(support_size, 1), inlier_scores], 0)\n                if self.use_inlier_latent\n                else torch.ones(support_size + len(inlier_scores), 1)\n            )  # [support_size + query_size, 1]\n            centroids = (\n                self.ema_weight * ((all_inliers_scores * all_targets).T @ all_features / (all_inliers_scores * all_targets).sum(0).unsqueeze(1))\n                + (1 - self.ema_weight) * centroids\n            )  # [num_classes, feature_dim]\n\n        logits_s = self.get_logits(centroids, support_features)\n        logits_q = self.get_logits(centroids, query_features)\n\n        if self.inference_steps == 0:\n            outlier_scores = (1.0 - (latent_targets * logits_q / self.lambda_s).sum(-1, keepdim=True).sigmoid())\n\n        return (logits_s.softmax(-1), logits_q.softmax(-1), outlier_scores)","metadata":{"execution":{"iopub.status.busy":"2024-01-01T02:51:19.572224Z","iopub.execute_input":"2024-01-01T02:51:19.572551Z","iopub.status.idle":"2024-01-01T02:51:19.637571Z","shell.execute_reply.started":"2024-01-01T02:51:19.572523Z","shell.execute_reply":"2024-01-01T02:51:19.636537Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Inference","metadata":{}},{"cell_type":"code","source":"IMAGES_PATH = '/kaggle/input/UBC-OCEAN/test_images'\nTHUMB_PATH = '/kaggle/input/UBC-OCEAN/test_thumbnails'\n#IMAGES_PATH = '/kaggle/input/UBC-OCEAN/train_images'\n#THUMB_PATH = '/kaggle/input/UBC-OCEAN/train_thumbnails'\nBATCH_SIZE = 1\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nk = 16\ninlier_thresh = 0.05","metadata":{"execution":{"iopub.status.busy":"2024-01-01T02:51:19.639061Z","iopub.execute_input":"2024-01-01T02:51:19.639942Z","iopub.status.idle":"2024-01-01T02:51:19.652515Z","shell.execute_reply.started":"2024-01-01T02:51:19.639902Z","shell.execute_reply":"2024-01-01T02:51:19.651499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = ['HGSC', 'LGSC', 'CC', 'EC', 'MC', 'Other']\nmetadata = pd.read_csv('/kaggle/input/UBC-OCEAN/test.csv')\n#bad_sample = [2906, 3191, 5851, 6281, 6363, 6898, 8279, 8713, 9183, 9254, 10252, 11263, 14051, 14401, 14617, 15231, 16209,\n#             22740, 25331, 27739, 29147, 29904, 32035, 34649, 34690, 34822, 36063, 36499, 38366, 42549, 44232, 44530, 48506,\n#             48550, 49587, 49872, 50962, 51032, 51128, 53377, 54007, 56117, 56351, 56947, 57162, 59031, 62476, 63165, 64111]\n#metadata = pd.read_csv('/kaggle/input/UBC-OCEAN/train.csv')\nmetadata['size'] = metadata['image_width'] * metadata['image_height']\nmetadata = metadata.sort_values(by=['image_id'], ascending=True)\nwsis = metadata.loc[metadata['size'] >= 30000000, 'image_id'].values.flatten().tolist()\n#query_sample = [x for x in query_sample if x not in bad_sample]\ntmas = metadata.loc[metadata['size'] < 30000000, 'image_id'].values.flatten().tolist()\n#query_sample = wsis + tmas\n#query_sample = tmas\n\ntransforms = A.Compose(\n        [A.Normalize(),\n            ToTensorV2()]\n    )\n\nwsi_loader = get_loaders(\n        IMAGES_PATH,\n        THUMB_PATH,\n        wsis,\n        BATCH_SIZE,\n        transforms,\n        num_workers=1\n    )\n\ntma_loader = get_loaders(\n        IMAGES_PATH,\n        THUMB_PATH,\n        tmas,\n        BATCH_SIZE,\n        transforms,\n        num_workers=4\n    )\n\nmodel = WSINet(backbone='convnextv2_tiny.fcmae_ft_in22k_in1k_384', in_channels=3, num_classes=5, \n               dropout=0.25, heads=8, pretrained=True, \n               pretrained_cfg_overlay=dict(file='/kaggle/input/convnextv2-tiny-fcmae-ft-in22k-in1k-384/model.safetensors')).to(device=device)\nload_checkpoint(torch.load('/kaggle/input/convnext-roformer-abmil/cluster352.pth.tar'), model)\n\nwith h5py.File('/kaggle/input/trained-features/cluster352.hdf5', 'r') as hf:\n    #train_ids = np.asarray(hf['image_ids'])\n    train_features = np.asarray(hf['features'])\n\nmodel.eval()\n\nquery_features = []\nquery_preds = []\n#query_preds = list(zip(wsis, np.repeat(labels[-1], len(wsis))))\nwith torch.no_grad():\n    wsi_pop = []\n    for index, (tiles, coords, image_id) in enumerate(wsi_loader):\n        tiles = tiles.squeeze(0)\n        coords = coords.squeeze(0)\n        print(tiles.size())\n        print(coords.size())\n        if tiles.size(0) == 1:\n            query_preds.append((image_id.item(), labels[-1]))\n            wsi_pop.append(index)\n            continue\n        X = tiles.float().to(device=device, non_blocking=True)\n        y_prob, pred, features = model(X, coords)\n        query_preds.append((image_id.item(), labels[pred.to(device='cpu').item()]))\n        query_features.append(features.view(-1).to(device='cpu'))\n        del X\n        del tiles\n    \n    if len(wsi_pop) > 0:\n        wsi_pop.reverse()\n        for i in wsi_pop:\n            wsis.pop(i)\n            \n    torch.cuda.empty_cache()\n    \n    tma_pop = []\n    for index, (tiles, coords, image_id) in enumerate(tma_loader):\n        tiles = tiles.squeeze(0)\n        coords = coords.squeeze(0)\n        print(tiles.size())\n        print(coords.size())\n        if tiles.size(0) == 1:\n            query_preds.append((image_id.item(), labels[-1]))\n            tma_pop.append(index)\n            continue\n        X = tiles.float().to(device=device, non_blocking=True)\n        y_prob, pred, features = model(X, coords)\n        query_preds.append((image_id.item(), labels[pred.to(device='cpu').item()]))\n        query_features.append(features.view(-1).to(device='cpu'))\n        del X\n        del tiles\n        \n    torch.cuda.empty_cache()\n    \n    if len(tma_pop) > 0:\n        tma_pop.reverse()\n        for i in tma_pop:\n            tmas.pop(i)\n\nquery_features = torch.stack(query_features) if len(query_features) > 0 else torch.stack([torch.tensor([])])\nsubmission = pd.DataFrame.from_records(query_preds, columns=['image_id', 'label'])\nquery_sample = wsis + tmas","metadata":{"execution":{"iopub.status.busy":"2024-01-01T02:56:34.810118Z","iopub.execute_input":"2024-01-01T02:56:34.810722Z","iopub.status.idle":"2024-01-01T02:57:10.356376Z","shell.execute_reply.started":"2024-01-01T02:56:34.81069Z","shell.execute_reply":"2024-01-01T02:57:10.355179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Outlier detection","metadata":{}},{"cell_type":"code","source":"knn_index = faiss.IndexFlatL2(train_features.shape[1])\nknn_index.add(train_features)\nD_t, _ = knn_index.search(train_features, k)\nscores_train = -D_t[:, -1]\nD_q, _ = knn_index.search(query_features, k)\nscores_query = -D_q[:, -1]\nfull_query = list(zip(query_sample, scores_query))\nscores_train.sort()\nfull_query = sorted(full_query, key=lambda x: x[1])\nthreshold = scores_train[round(inlier_thresh * scores_train.shape[0])]\noutliers = [x[0] for x in full_query if x[1] < threshold]\nsubmission['label'] = np.where(submission['image_id'].isin(outliers), 'Other', submission['label'])\nsubmission = submission.sort_values(by=['image_id'])","metadata":{"execution":{"iopub.status.busy":"2024-01-01T02:57:10.358798Z","iopub.execute_input":"2024-01-01T02:57:10.359592Z","iopub.status.idle":"2024-01-01T02:57:10.473836Z","shell.execute_reply.started":"2024-01-01T02:57:10.359552Z","shell.execute_reply":"2024-01-01T02:57:10.472971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Submission","metadata":{}},{"cell_type":"code","source":"submission.to_csv('submission.csv', index=False)\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2024-01-01T02:57:10.475081Z","iopub.execute_input":"2024-01-01T02:57:10.475416Z","iopub.status.idle":"2024-01-01T02:57:10.493591Z","shell.execute_reply.started":"2024-01-01T02:57:10.475388Z","shell.execute_reply":"2024-01-01T02:57:10.492634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2024-01-01T02:57:10.495307Z","iopub.execute_input":"2024-01-01T02:57:10.495597Z","iopub.status.idle":"2024-01-01T02:57:10.503514Z","shell.execute_reply.started":"2024-01-01T02:57:10.495572Z","shell.execute_reply":"2024-01-01T02:57:10.502536Z"},"trusted":true},"execution_count":null,"outputs":[]}]}