{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":49349,"databundleVersionId":5447706,"sourceType":"competition"},{"sourceId":5868754,"sourceType":"datasetVersion","datasetId":3374238}],"dockerImageVersionId":30646,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\nThis is part of my attempt to implement (and learn from) the [2nd place solution at IMC2023](https://www.kaggle.com/competitions/image-matching-challenge-2023/discussion/416873).  \n\nContributions:  \n* SuperPoint is used with `conf_thresh = 0.005` and with input resolutions `[1088, 1280, 1376]`.\n* SIFT was used to show comparison.\n* My takeaways are: \n    * SIFT might be a beter keypoint detector than SuperPoint after all.  \n    * Increasing input image resolution for SuperPoint doesn't necessarily mean a better number or quality of keypoints.  \n    * SuperPoint works better in the dark compared to SIFT.\n\nWas not expecting to reach these conclusions when I started the notebook.  \n\n\nPrevious Notebook:  \n[rotation-correction-imc2023-2nd-position-solution](https://www.kaggle.com/code/mukit0/rotation-correction-imc2023-2nd-position-solution#Rotation-Correction) the rotation correction used by the 2nd place solution.  ","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport cv2\nimport torch\nfrom torchvision.transforms import functional as F\nimport time\nfrom matplotlib import pyplot as plt\nfrom tqdm import tqdm\n!pip install mediapy\nimport mediapy","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:14:25.940879Z","iopub.execute_input":"2024-02-27T03:14:25.941248Z","iopub.status.idle":"2024-02-27T03:14:35.684458Z","shell.execute_reply.started":"2024-02-27T03:14:25.94122Z","shell.execute_reply":"2024-02-27T03:14:35.683643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setup SuperPoint API and Configuration","metadata":{}},{"cell_type":"markdown","source":"## Setup Class Architecture and Frontend\nThese are boilerplate code we can find from SuperPoint github repo. https://github.com/magicleap/SuperPointPretrainedNetwork","metadata":{}},{"cell_type":"code","source":"class SuperPointNet(torch.nn.Module):\n  \"\"\" Pytorch definition of SuperPoint Network. \"\"\"\n  def __init__(self):\n    super(SuperPointNet, self).__init__()\n    self.relu = torch.nn.ReLU(inplace=True)\n    self.pool = torch.nn.MaxPool2d(kernel_size=2, stride=2)\n    c1, c2, c3, c4, c5, d1 = 64, 64, 128, 128, 256, 256\n    # Shared Encoder.\n    self.conv1a = torch.nn.Conv2d(1, c1, kernel_size=3, stride=1, padding=1)\n    self.conv1b = torch.nn.Conv2d(c1, c1, kernel_size=3, stride=1, padding=1)\n    self.conv2a = torch.nn.Conv2d(c1, c2, kernel_size=3, stride=1, padding=1)\n    self.conv2b = torch.nn.Conv2d(c2, c2, kernel_size=3, stride=1, padding=1)\n    self.conv3a = torch.nn.Conv2d(c2, c3, kernel_size=3, stride=1, padding=1)\n    self.conv3b = torch.nn.Conv2d(c3, c3, kernel_size=3, stride=1, padding=1)\n    self.conv4a = torch.nn.Conv2d(c3, c4, kernel_size=3, stride=1, padding=1)\n    self.conv4b = torch.nn.Conv2d(c4, c4, kernel_size=3, stride=1, padding=1)\n    # Detector Head.\n    self.convPa = torch.nn.Conv2d(c4, c5, kernel_size=3, stride=1, padding=1)\n    self.convPb = torch.nn.Conv2d(c5, 65, kernel_size=1, stride=1, padding=0)\n    # Descriptor Head.\n    self.convDa = torch.nn.Conv2d(c4, c5, kernel_size=3, stride=1, padding=1)\n    self.convDb = torch.nn.Conv2d(c5, d1, kernel_size=1, stride=1, padding=0)\n\n  def forward(self, x):\n    \"\"\" Forward pass that jointly computes unprocessed point and descriptor\n    tensors.\n    Input\n      x: Image pytorch tensor shaped N x 1 x H x W.\n    Output\n      semi: Output point pytorch tensor shaped N x 65 x H/8 x W/8.\n      desc: Output descriptor pytorch tensor shaped N x 256 x H/8 x W/8.\n    \"\"\"\n    # Shared Encoder.\n    x = self.relu(self.conv1a(x))\n    x = self.relu(self.conv1b(x))\n    x = self.pool(x)\n    x = self.relu(self.conv2a(x))\n    x = self.relu(self.conv2b(x))\n    x = self.pool(x)\n    x = self.relu(self.conv3a(x))\n    x = self.relu(self.conv3b(x))\n    x = self.pool(x)\n    x = self.relu(self.conv4a(x))\n    x = self.relu(self.conv4b(x))\n    # Detector Head.\n    cPa = self.relu(self.convPa(x))\n    semi = self.convPb(cPa)\n    # Descriptor Head.\n    cDa = self.relu(self.convDa(x))\n    desc = self.convDb(cDa)\n    dn = torch.norm(desc, p=2, dim=1) # Compute the norm.\n    desc = desc.div(torch.unsqueeze(dn, 1)) # Divide by norm to normalize.\n    return semi, desc","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:14:35.686496Z","iopub.execute_input":"2024-02-27T03:14:35.686847Z","iopub.status.idle":"2024-02-27T03:14:35.701603Z","shell.execute_reply.started":"2024-02-27T03:14:35.686814Z","shell.execute_reply":"2024-02-27T03:14:35.70044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SuperPointFrontend(object):\n  \"\"\" Wrapper around pytorch net to help with pre and post image processing. \"\"\"\n  def __init__(self, weights_path, nms_dist, conf_thresh, nn_thresh,\n               cuda=False):\n    self.name = 'SuperPoint'\n    self.cuda = cuda\n    self.nms_dist = nms_dist\n    self.conf_thresh = conf_thresh\n    self.nn_thresh = nn_thresh # L2 descriptor distance for good match.\n    self.cell = 8 # Size of each output cell. Keep this fixed.\n    self.border_remove = 4 # Remove points this close to the border.\n\n    # Load the network in inference mode.\n    self.net = SuperPointNet()\n    if cuda:\n      # Train on GPU, deploy on GPU.\n      self.net.load_state_dict(torch.load(weights_path))\n      self.net = self.net.cuda()\n    else:\n      # Train on GPU, deploy on CPU.\n      self.net.load_state_dict(torch.load(weights_path,\n                               map_location=lambda storage, loc: storage))\n    self.net.eval()\n\n  def nms_fast(self, in_corners, H, W, dist_thresh):\n    \"\"\"\n    Run a faster approximate Non-Max-Suppression on numpy corners shaped:\n      3xN [x_i,y_i,conf_i]^T\n  \n    Algo summary: Create a grid sized HxW. Assign each corner location a 1, rest\n    are zeros. Iterate through all the 1's and convert them either to -1 or 0.\n    Suppress points by setting nearby values to 0.\n  \n    Grid Value Legend:\n    -1 : Kept.\n     0 : Empty or suppressed.\n     1 : To be processed (converted to either kept or supressed).\n  \n    NOTE: The NMS first rounds points to integers, so NMS distance might not\n    be exactly dist_thresh. It also assumes points are within image boundaries.\n  \n    Inputs\n      in_corners - 3xN numpy array with corners [x_i, y_i, confidence_i]^T.\n      H - Image height.\n      W - Image width.\n      dist_thresh - Distance to suppress, measured as an infinty norm distance.\n    Returns\n      nmsed_corners - 3xN numpy matrix with surviving corners.\n      nmsed_inds - N length numpy vector with surviving corner indices.\n    \"\"\"\n    grid = np.zeros((H, W)).astype(int) # Track NMS data.\n    inds = np.zeros((H, W)).astype(int) # Store indices of points.\n    # Sort by confidence and round to nearest int.\n    inds1 = np.argsort(-in_corners[2,:])\n    corners = in_corners[:,inds1]\n    rcorners = corners[:2,:].round().astype(int) # Rounded corners.\n    # Check for edge case of 0 or 1 corners.\n    if rcorners.shape[1] == 0:\n      return np.zeros((3,0)).astype(int), np.zeros(0).astype(int)\n    if rcorners.shape[1] == 1:\n      out = np.vstack((rcorners, in_corners[2])).reshape(3,1)\n      return out, np.zeros((1)).astype(int)\n    # Initialize the grid.\n    for i, rc in enumerate(rcorners.T):\n      grid[rcorners[1,i], rcorners[0,i]] = 1\n      inds[rcorners[1,i], rcorners[0,i]] = i\n    # Pad the border of the grid, so that we can NMS points near the border.\n    pad = dist_thresh\n    grid = np.pad(grid, ((pad,pad), (pad,pad)), mode='constant')\n    # Iterate through points, highest to lowest conf, suppress neighborhood.\n    count = 0\n    for i, rc in enumerate(rcorners.T):\n      # Account for top and left padding.\n      pt = (rc[0]+pad, rc[1]+pad)\n      if grid[pt[1], pt[0]] == 1: # If not yet suppressed.\n        grid[pt[1]-pad:pt[1]+pad+1, pt[0]-pad:pt[0]+pad+1] = 0\n        grid[pt[1], pt[0]] = -1\n        count += 1\n    # Get all surviving -1's and return sorted array of remaining corners.\n    keepy, keepx = np.where(grid==-1)\n    keepy, keepx = keepy - pad, keepx - pad\n    inds_keep = inds[keepy, keepx]\n    out = corners[:, inds_keep]\n    values = out[-1, :]\n    inds2 = np.argsort(-values)\n    out = out[:, inds2]\n    out_inds = inds1[inds_keep[inds2]]\n    return out, out_inds\n\n  def run(self, img):\n    \"\"\" Process a numpy image to extract points and descriptors.\n    Input\n      img - HxW numpy float32 input image in range [0,1].\n    Output\n      corners - 3xN numpy array with corners [x_i, y_i, confidence_i]^T.\n      desc - 256xN numpy array of corresponding unit normalized descriptors.\n      heatmap - HxW numpy heatmap in range [0,1] of point confidences.\n      \"\"\"\n    assert img.ndim == 2, 'Image must be grayscale.'\n    assert img.dtype == np.float32, 'Image must be float32.'\n    H, W = img.shape[0], img.shape[1]\n    inp = img.copy()\n    inp = (inp.reshape(1, H, W))\n    inp = torch.from_numpy(inp)\n    inp = torch.autograd.Variable(inp).view(1, 1, H, W)\n    if self.cuda:\n      inp = inp.cuda()\n    # Forward pass of network.\n    outs = self.net.forward(inp)\n    semi, coarse_desc = outs[0], outs[1]\n    # Convert pytorch -> numpy.\n    semi = semi.data.cpu().numpy().squeeze()\n    # --- Process points.\n    dense = np.exp(semi) # Softmax.\n    dense = dense / (np.sum(dense, axis=0)+.00001) # Should sum to 1.\n    # Remove dustbin.\n    nodust = dense[:-1, :, :]\n    # Reshape to get full resolution heatmap.\n    Hc = int(H / self.cell)\n    Wc = int(W / self.cell)\n    nodust = nodust.transpose(1, 2, 0)\n    heatmap = np.reshape(nodust, [Hc, Wc, self.cell, self.cell])\n    heatmap = np.transpose(heatmap, [0, 2, 1, 3])\n    heatmap = np.reshape(heatmap, [Hc*self.cell, Wc*self.cell])\n    xs, ys = np.where(heatmap >= self.conf_thresh) # Confidence threshold.\n    if len(xs) == 0:\n      return np.zeros((3, 0)), None, None\n    pts = np.zeros((3, len(xs))) # Populate point data sized 3xN.\n    pts[0, :] = ys\n    pts[1, :] = xs\n    pts[2, :] = heatmap[xs, ys]\n    pts, _ = self.nms_fast(pts, H, W, dist_thresh=self.nms_dist) # Apply NMS.\n    inds = np.argsort(pts[2,:])\n    pts = pts[:,inds[::-1]] # Sort by confidence.\n    # Remove points along border.\n    bord = self.border_remove\n    toremoveW = np.logical_or(pts[0, :] < bord, pts[0, :] >= (W-bord))\n    toremoveH = np.logical_or(pts[1, :] < bord, pts[1, :] >= (H-bord))\n    toremove = np.logical_or(toremoveW, toremoveH)\n    pts = pts[:, ~toremove]\n    # --- Process descriptor.\n    D = coarse_desc.shape[1]\n    if pts.shape[1] == 0:\n      desc = np.zeros((D, 0))\n    else:\n      # Interpolate into descriptor map using 2D point locations.\n      samp_pts = torch.from_numpy(pts[:2, :].copy())\n      samp_pts[0, :] = (samp_pts[0, :] / (float(W)/2.)) - 1.\n      samp_pts[1, :] = (samp_pts[1, :] / (float(H)/2.)) - 1.\n      samp_pts = samp_pts.transpose(0, 1).contiguous()\n      samp_pts = samp_pts.view(1, 1, -1, 2)\n      samp_pts = samp_pts.float()\n      if self.cuda:\n        samp_pts = samp_pts.cuda()\n      desc = torch.nn.functional.grid_sample(coarse_desc, samp_pts)\n      desc = desc.data.cpu().numpy().reshape(D, -1)\n      desc /= np.linalg.norm(desc, axis=0)[np.newaxis, :]\n    return pts, desc, heatmap","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:14:35.703094Z","iopub.execute_input":"2024-02-27T03:14:35.703473Z","iopub.status.idle":"2024-02-27T03:14:35.729924Z","shell.execute_reply.started":"2024-02-27T03:14:35.703438Z","shell.execute_reply":"2024-02-27T03:14:35.729065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Configure SuperPoint\nIn their [solution](https://www.kaggle.com/competitions/image-matching-challenge-2023/discussion/416873), the authors used a confidence threshold of 0.005. ","metadata":{}},{"cell_type":"code","source":"weights_path = '/kaggle/input/superpoint-magicleap/superpoint_v1.pth'\nfrontend = SuperPointFrontend(weights_path, nms_dist=4, conf_thresh=0.005, nn_thresh=0.7, cuda=False)","metadata":{"_kg_hide-input":false,"_kg_hide-output":false,"execution":{"iopub.status.busy":"2024-02-27T03:14:35.731773Z","iopub.execute_input":"2024-02-27T03:14:35.732057Z","iopub.status.idle":"2024-02-27T03:14:35.770974Z","shell.execute_reply.started":"2024-02-27T03:14:35.732016Z","shell.execute_reply":"2024-02-27T03:14:35.770171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load and Show Images From IMC2023","metadata":{}},{"cell_type":"code","source":"from glob import glob\nfrom pprint import pprint\nimagesDirList = glob('/kaggle/input/image-matching-challenge-2023/train/*/*/images')\npprint(imagesDirList)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:14:35.772063Z","iopub.execute_input":"2024-02-27T03:14:35.772325Z","iopub.status.idle":"2024-02-27T03:14:35.802676Z","shell.execute_reply.started":"2024-02-27T03:14:35.772302Z","shell.execute_reply":"2024-02-27T03:14:35.801708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Please feel free to change the `dirIdx` to choose any other dataset.","metadata":{}},{"cell_type":"code","source":"dirIdx = 5\nsrc = imagesDirList[dirIdx]\nimages = [cv2.cvtColor(cv2.imread(im), cv2.COLOR_BGR2RGB) for im in glob(f'{src}/*')]\nprint(f'Found {len(images)} images in {src}')","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:14:35.805116Z","iopub.execute_input":"2024-02-27T03:14:35.805608Z","iopub.status.idle":"2024-02-27T03:14:46.597216Z","shell.execute_reply.started":"2024-02-27T03:14:35.805577Z","shell.execute_reply":"2024-02-27T03:14:46.596292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maxImageCount = 25\nmaxImageCount = min(maxImageCount, len(images))\nprint(f'Showing {maxImageCount} image from {src}')\nmediapy.show_images(images[:maxImageCount], height=300, columns=5)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:14:46.598503Z","iopub.execute_input":"2024-02-27T03:14:46.598794Z","iopub.status.idle":"2024-02-27T03:14:48.174062Z","shell.execute_reply.started":"2024-02-27T03:14:46.59877Z","shell.execute_reply":"2024-02-27T03:14:48.172773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Detect SuperPoint Keypoints for Multiple Resolutions","metadata":{}},{"cell_type":"markdown","source":"From the authors:  \n> TTA. Ensemble of matches extracted from images at different scales. In our local experiments, the best results are achieved with a combination of [1088, 1280, 1376]. We used np.concatenate to join matches from different models that was pretty common for the last IMC22 competition.\n\nTo detect keypoints on different input shapes, I implemented the following functions. `resizeToResolution` resizes an input image to the desired resolution. Once the keypoints are detected by SuperPoint, we need to map these keypoints back to the original input image's dimensions. The function `rescaleKeypoints` helps with that. I also added a simple sanity check where I see if unit vectors are mapped back correctly.  \n\nReferences:  \nhttps://www.kaggle.com/competitions/image-matching-challenge-2023/discussion/416873  \n[IMC 2022-kornia : Score 0.725](https://www.kaggle.com/code/cbeaud/imc-2022-kornia-score-0-725)  ","metadata":{}},{"cell_type":"code","source":"def resizeToResolution(img, resolution):\n    scale = resolution / max(img.shape[:2])\n    h = int(img.shape[0] * scale)\n    w = int(img.shape[1] * scale)\n    return cv2.resize(img, (w, h))\n\ndef rescaleKeypoints(keypoints, resolution, resizeToShape):\n    scale = resolution / max(resizeToShape)\n    return keypoints/scale\n\n# Sanity checks\nresolution = 1088\nimgIdx = 0\nscaleFactor = resolution / max(images[imgIdx].shape[:2])\nprint(f\"{images[imgIdx].shape[:2]} needs to be scaled by {scaleFactor} to be of resolution {resolution}\")\nrescaledImage = resizeToResolution(images[imgIdx], resolution)\nprint(f\"Rescaled image shape: {rescaledImage.shape}\")\n\ndummyKeypoints = np.array([[0,1], [1, 0]], dtype='float32')\nprint(f\"Image was scaled by a factor of {scaleFactor} so keypoints should scaled back to {1/scaleFactor}\")\nrescaledKeyPoints = rescaleKeypoints(dummyKeypoints, resolution, images[imgIdx].shape[:2])\nprint(f\"Input keypoints:\\n{dummyKeypoints}\")\nprint(f\"Recaled back Keypoints:\\n{rescaledKeyPoints}\")","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:14:48.175111Z","iopub.execute_input":"2024-02-27T03:14:48.175394Z","iopub.status.idle":"2024-02-27T03:14:48.186149Z","shell.execute_reply.started":"2024-02-27T03:14:48.175367Z","shell.execute_reply":"2024-02-27T03:14:48.185315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"resolutionList = [1088, 1280, 1376]\nresolutionList.sort()\nkeypointsPerResolution = {resolution: [] for resolution in resolutionList}\nprint(keypointsPerResolution)\n\nstartTime = time.time()\n\nfor imgIdx in tqdm(range(maxImageCount)):\n    imgGray = cv2.cvtColor(images[imgIdx], cv2.COLOR_RGB2GRAY)  # Note: images were already converted to RGB before.\n    imgGrayNormalized = imgGray.astype('float32') / 255.0  # Normalizing to keep value in range 0-1 as required by SuperPointFrontEnd\n    \n    for resolution in resolutionList:\n        # ------ Apply SuperPoint and store output ---------\n        imgGrayResized = resizeToResolution(imgGrayNormalized, resolution)\n        corners, descriptors, heatmap = frontend.run(imgGrayResized)\n        keypointsSuperPoint = corners[:2, :].T\n        keypointsSuperPoint = rescaleKeypoints(keypointsSuperPoint, resolution, imgGray.shape[:2])\n        keypointsPerResolution[resolution].append(keypointsSuperPoint)\n\nprint(f\"Took: {time.time() - startTime} seconds.\")","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:14:48.187169Z","iopub.execute_input":"2024-02-27T03:14:48.187467Z","iopub.status.idle":"2024-02-27T03:16:58.398811Z","shell.execute_reply.started":"2024-02-27T03:14:48.187443Z","shell.execute_reply":"2024-02-27T03:16:58.397879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Detect SIFT Keypoints","metadata":{}},{"cell_type":"code","source":"featuresXYSIFT = []\nsift = cv2.SIFT_create()\n\nstartTime = time.time()\nfor imgIdx in tqdm(range(maxImageCount)):\n    img = images[imgIdx].copy()\n    imgGray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)  # Note: images were already converted to RGB before.\n    \n    # ------ Apply SIFT and store output ---------\n    keypointsSIFT = sift.detect(imgGray, None)\n    featuresXYSIFT.append(np.array(list(map(lambda kp: kp.pt, keypointsSIFT))))\n\nprint(f\"Took: {time.time() - startTime} seconds.\")","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:16:58.402014Z","iopub.execute_input":"2024-02-27T03:16:58.403145Z","iopub.status.idle":"2024-02-27T03:17:01.008995Z","shell.execute_reply.started":"2024-02-27T03:16:58.403112Z","shell.execute_reply":"2024-02-27T03:17:01.008163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Compare Keypoint Detections","metadata":{}},{"cell_type":"markdown","source":"## Plot Count of Keypoints","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(figsize = (10, 10))\nsuperPointPlot1088, = ax.plot([len(kps) for kps in keypointsPerResolution[1088]], 'ro-')\nsuperPointPlot1280, = ax.plot([len(kps) for kps in keypointsPerResolution[1280]], 'go-')\nsuperPointPlot1376, = ax.plot([len(kps) for kps in keypointsPerResolution[1376]], 'bo-')\nsiftPlot, = ax.plot([len(kps) for kps in featuresXYSIFT], 'ko-')\nax.set_title(\"Comparison of counts of keypoints obtained by different methods.\")\nax.set_xlabel(\"ImgIdx\")\nax.set_ylabel(\"Count of keypoints.\")\nax.legend(\n    (superPointPlot1088, superPointPlot1280, superPointPlot1376, siftPlot),\n    (\"SP 1088\", \"SP 1280\", \"SP 1376\", \"SIFT\"),\n    loc='upper right',\n    shadow=True)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:17:01.010079Z","iopub.execute_input":"2024-02-27T03:17:01.010383Z","iopub.status.idle":"2024-02-27T03:17:01.328294Z","shell.execute_reply.started":"2024-02-27T03:17:01.010355Z","shell.execute_reply":"2024-02-27T03:17:01.327103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"From the above plots, (for this dataset and selection of images) SIFT seems to produce the most number of keypoints consistantly. For SuperPoint, the resolution 1280 seems often outperform resolution 1376. The relationship between resolution and number of keypoints produced by SuperPoint doesn't seem to be very linear.","metadata":{}},{"cell_type":"markdown","source":"## Draw Keypoints on Image for Visual Comparison","metadata":{}},{"cell_type":"code","source":"def drawKeypoints(keypoints, displayImage, color=(255, 0, 0)):\n    for kp in keypoints:\n        x, y = kp.astype('int32')\n        cv2.circle(displayImage, (x, y), 2, color, -1)\n    return displayImage","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:17:01.329438Z","iopub.execute_input":"2024-02-27T03:17:01.329703Z","iopub.status.idle":"2024-02-27T03:17:01.335008Z","shell.execute_reply.started":"2024-02-27T03:17:01.329679Z","shell.execute_reply":"2024-02-27T03:17:01.333959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Observations:\nKeypoints from each method is drawn below. It looks like the image-scaling and then scaling-back the keypoints worked.  \nSuperPoint often detects kepoints in the sky or plain textureless spaces. These are not tractable features from image to image. As input resolution increased the number of these \"unnecessary\" points only increased. I am not sure if increasing such number of points by increasing the resolution is really worth it when we will perform descriptor matching.\n\nIn comparison, SIFT seems to cover all important textured places with keypoint detections. It didn't detect keypoints in textureless spaces nearly as much as SuperPoint. From these experiments, I still feel more confident about SIFT than SuperPoint when comes to detecting tractable keypoints.","metadata":{}},{"cell_type":"code","source":"from torchvision.io import read_image as T_read_image\nfrom torchvision.io import ImageReadMode\nfrom torchvision import transforms as T\nfrom check_orientation.pre_trained_models import create_model\n\n# General utilities\nimport os\nfrom tqdm import tqdm\nfrom time import time\nfrom fastprogress import progress_bar\nimport gc\nimport numpy as np\nimport h5py\nfrom IPython.display import clear_output\nfrom collections import defaultdict\nfrom copy import deepcopy\nimport pandas as pd\nfrom dataclasses import dataclass\n\n# CV/ML\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport kornia as K\nimport kornia.feature as KF\nfrom PIL import Image\nimport timm\nfrom timm.data import resolve_data_config\nfrom timm.data.transforms_factory import create_transform\nfrom torchvision import transforms\n\n# We will use ViT global descriptor to get matching shortlists.\ndef get_global_desc(fnames, model, model2,\n                    device =  device):\n    scalepic=600\n    model = model.eval()\n    model = model.to(device)\n    model2 = model2.eval()\n    model2 = model2.to(device)\n    config = resolve_data_config({}, model=model)\n    transform = transforms.Compose([\n                transforms.Resize(scalepic, interpolation=transforms.InterpolationMode.BICUBIC),\n                transforms.CenterCrop(scalepic),\n                transforms.ToTensor(),\n                transforms.Normalize([0.4850, 0.4560, 0.4060], [0.2290, 0.2240, 0.2250]),])#create_transform(**config)\n    global_descs_convnext=[]\n    for i, img_fname_full in tqdm(enumerate(fnames),total= len(fnames)):\n        key = os.path.splitext(os.path.basename(img_fname_full))[0]\n        img = Image.open(img_fname_full).convert('RGB')\n        timg = transform(img).unsqueeze(0).to(device)\n        with torch.no_grad():\n            desc = model.forward_features(timg.to(device)).mean(dim=(-1,2))\n            desc2 = model2.forward_features(timg.to(device)).mean(dim=(-1,2))\n            desc = desc.view(1, -1)\n            desc2 = desc2.view(1, -1)\n            desc_norm = torch.cat([desc, desc2], dim=-1)\n            desc_norm = F.normalize(desc_norm, dim=1, p=2)\n        global_descs_convnext.append(desc_norm.detach().cpu())\n    global_descs_all = torch.cat(global_descs_convnext, dim=0)\n    return global_descs_all\n\n\ndef get_img_pairs_exhaustive(img_fnames):\n    index_pairs = []\n    for i in range(len(img_fnames)):\n        for j in range(i+1, len(img_fnames)):\n            index_pairs.append((i,j))\n    return index_pairs\ndef get_image_pairs_shortlist(fnames,\n                            exhaustive_if_less = 20,\n                            nneighbor=60,\n                            device=torch.device('cpu'),\n                            th=0.1):\n    num_imgs = len(fnames)\n    if num_imgs <= exhaustive_if_less or num_imgs <= nneighbor :\n        return get_img_pairs_exhaustive(fnames)\n\n    model = timm.create_model('tf_efficientnet_b6',\n                            checkpoint_path='/kaggle/input/tf-efficientnet/pytorch/tf-efficientnet-b6/1/tf_efficientnet_b6_aa-80ba17e4.pth')\n    model.eval()\n    model2 = timm.create_model('tf_efficientnet_b7',\n                            checkpoint_path='/kaggle/input/tf-efficientnet/pytorch/tf-efficientnet-b7/1/tf_efficientnet_b7_ra-6c08e654.pth')                  \n    model2.eval()\n    descs = get_global_desc(fnames, model, model2, device=device)\n    del model, model2\n    torch.cuda.empty_cache()\n    import gc\n    gc.collect()\n    dm = torch.einsum('bi,ki->bk', descs, descs).detach().cpu() \n    value, index = torch.topk(dm, k=nneighbor, dim=1)\n\n    matching_list = []\n    for i in range(num_imgs-1):\n        for t in index[i][value[i]>th]:\n            if t == i:\n                continue\n            matching_list.append(tuple(sorted((i, t.item()))))\n    matching_list = sorted(list(set(matching_list)))\n    return matching_list\n\ndef load_torch_image(fname, device=torch.device('cpu')):\n    img = K.io.load_image(fname, K.io.ImageLoadType.RGB32, device=device)[None, ...]\n    return img\n\ndef convert_coord(r, w, h, rotk):\n    if rotk == 0:\n        return r\n    elif rotk == 1:\n        rx = w-1-r[:, 1]\n        ry = r[:, 0]\n        return torch.concat([rx[None], ry[None]], dim=0).T\n    elif rotk == 2:\n        rx = w-1-r[:, 0]\n        ry = h-1-r[:, 1]\n        return torch.concat([rx[None], ry[None]], dim=0).T\n    elif rotk == 3:\n        rx = r[:, 1]\n        ry = h-1-r[:, 0]\n        return torch.concat([rx[None], ry[None]], dim=0).T\n\ndef detect_common(img_fnames,\n                model_name,\n                rots,\n                file_keypoints,\n                feature_dir = '.featureout',\n                num_features = 4096,\n                resize_to = 1024,\n                detection_threshold = 0.01,\n                device=torch.device('cpu'),\n                min_matches=15,verbose=True\n                ):\n    if not os.path.isdir(feature_dir):\n        os.makedirs(feature_dir)\n    dict_model = {\n        \"aliked\" : ALIKED,\n        \"superpoint\" : SuperPoint,\n        \"doghardnet\" : DoGHardNet,\n        \"disk\" : DISK,\n        \"sift\" : SIFT,\n    }\n    extractor_class = dict_model[model_name]\n\n    dtype = torch.float32 # ALIKED has issues with float16\n    extractor = extractor_class(\n        max_num_keypoints=num_features, detection_threshold=detection_threshold, resize=resize_to\n    ).eval().to(device, dtype)\n    rot=True\n    dict_kpts_cuda = {}\n    dict_descs_cuda = {}\n    for (img_path, rot_k) in zip(img_fnames, rots):\n        img_fname = img_path.split('/')[-1]\n        key = img_fname\n        with torch.inference_mode():\n            image0 = load_torch_image(img_path, device=device).to(dtype)\n            h, w = image0.shape[2], image0.shape[3]\n            image1 = torch.rot90(image0, rot_k, [2, 3])\n            feats0 = extractor.extract(image1)  # auto-resize the image, disable with resize=None\n            kpts = feats0['keypoints'].reshape(-1, 2).detach()\n            descs = feats0['descriptors'].reshape(len(kpts), -1).detach()\n            if rot==True:\n                kpts = convert_coord(kpts, w, h, rot_k)\n            dict_kpts_cuda[f\"{key}\"] = kpts\n            dict_descs_cuda[f\"{key}\"] = descs\n            print(f\"{model_name} > rot_k={rot_k}, kpts.shape={kpts.shape}, descs.shape={descs.shape}\")\n    del extractor\n    gc.collect()\n\n    #####################################################\n    # Matching keypoints\n    #####################################################\n    lg_matcher = KF.LightGlueMatcher(model_name, {\"width_confidence\": -1,\n                                            \"depth_confidence\": -1,\n                                            \"mp\": True if 'cuda' in str(device) else False}).eval().to(device)\n\n    cnt_pairs = 0\n    with h5py.File(file_keypoints, mode='w') as f_match:\n        for pair_idx in tqdm(index_pairs):\n            idx1, idx2 = pair_idx\n            fname1, fname2 = img_fnames[idx1], img_fnames[idx2]\n\n            key1, key2 = fname1.split('/')[-1], fname2.split('/')[-1]\n\n\n            kp1 = dict_kpts_cuda[key1]\n            kp2 = dict_kpts_cuda[key2]\n\n            desc1 = dict_descs_cuda[key1]\n            desc2 = dict_descs_cuda[key2]\n            with torch.inference_mode():\n                dists, idxs = lg_matcher(desc1,\n                                    desc2,\n                                    KF.laf_from_center_scale_ori(kp1[None]),\n                                    KF.laf_from_center_scale_ori(kp2[None]))\n            if len(idxs)  == 0:\n                continue\n            n_matches = len(idxs)\n            kp1 = kp1[idxs[:,0], :].cpu().numpy().reshape(-1, 2).astype(np.float32)\n            kp2 = kp2[idxs[:,1], :].cpu().numpy().reshape(-1, 2).astype(np.float32)\n            group  = f_match.require_group(key1)\n            if n_matches >= min_matches:\n                group.create_dataset(key2, data=np.concatenate([kp1, kp2], axis=1))\n                cnt_pairs+=1\n                print (f'{model_name}> {key1}-{key2}: {n_matches} matches @ {cnt_pairs}th pair({model_name}+lightglue)')            \n            else:\n                print (f'{model_name}> {key1}-{key2}: {n_matches} matches --> skipped')\n    del lg_matcher\n    torch.cuda.empty_cache()\n    gc.collect()\n    return\n\ndef detect_lightglue_common(\n    img_fnames, model_name, index_pairs, feature_dir, device, file_keypoints, rots,\n    resize_to=1024,\n    detection_threshold=0.01, \n    num_features=4096, \n    min_matches=15,\n):\n    t=time()\n    detect_common(\n        img_fnames, model_name, rots, file_keypoints, feature_dir, \n        resize_to=resize_to,\n        num_features=num_features, \n        detection_threshold=detection_threshold, \n        device=device,\n        min_matches=min_matches,\n    )\n    gc.collect()\n    t=time() -t \n    print(f'Features matched in  {t:.4f} sec ({model_name}+LightGlue)')\n    return t\n\n\nimport sys\nsys.path.append(\"../input/super-glue-pretrained-network\")\nfrom models.matching import Matching\nfrom models.superpoint import SuperPoint as SG_SuperPoint\nfrom models.superglue import SuperGlue\nfrom models.utils import (compute_pose_error, compute_epipolar_error,\n                        estimate_pose, make_matching_plot,\n                        error_colormap, AverageTimer, pose_auc, read_image,\n                        process_resize, frame2tensor,\n                        rotate_intrinsics, rotate_pose_inplane,\n                        scale_intrinsics)\n\nfrom torch.nn import functional as torchF  # For resizing tensor\n\n# Preprocess\ndef sg_read_image(image, device, resize):\n    w, h = image.shape[1], image.shape[0]\n    w_new, h_new = process_resize(w, h, [resize,])\n\n    unit_shape = 8\n    w_new = w_new // unit_shape * unit_shape\n    h_new = h_new // unit_shape * unit_shape\n\n    scales = (float(w) / float(w_new), float(h) / float(h_new))\n    image = cv2.resize(image.astype('float32'), (w_new, h_new))\n\n    inp = frame2tensor(image, \"cpu\")\n    return image, inp, scales, (h, w)\n\nclass SGDataset(Dataset):\n    def __init__(self, img_fnames, resize_to, device):\n        self.img_fnames = img_fnames\n        self.resize_to = resize_to\n        self.device = device\n\n    def __len__(self):\n        return len(self.img_fnames)\n\n    def __getitem__(self, idx):\n        fname = self.img_fnames[idx]\n        im = cv2.imread(fname, cv2.IMREAD_GRAYSCALE)\n        _, image, scale, ori_shape = sg_read_image(im, self.device, self.resize_to)\n        return image, torch.tensor([idx]), torch.tensor(ori_shape)\n\ndef get_superglue_dataloader(img_fnames, resize_to, device, batch_size=1):\n    dataset = SGDataset(img_fnames, resize_to, device)\n    dataloader = DataLoader(\n        dataset=dataset,\n        shuffle=False,\n        batch_size=batch_size,\n        pin_memory=True,\n        num_workers=2,\n        drop_last=False\n    )\n    return dataloader\n\ndef detect_superglue(\n    img_fnames, index_pairs, feature_dir, device, sg_config, file_keypoints, file_keypoints_crop,\n    resize_to=750, min_matches=15\n):    \n    t=time()\n\n    fnames1, fnames2, idxs1, idxs2 = [], [], [], []\n    for pair_idx in progress_bar(index_pairs):\n        idx1, idx2 = pair_idx\n        fname1, fname2 = img_fnames[idx1], img_fnames[idx2]\n        fnames1.append(fname1)\n        fnames2.append(fname2)\n        idxs1.append(idx1)\n        idxs2.append(idx2)\n\n    dataloader = get_superglue_dataloader(img_fnames=img_fnames, resize_to=1680, device=device)\n\n    #####################################################\n    # Extract keypoints and descriptions\n    #####################################################\n    superpoint = SG_SuperPoint(sg_config[\"superpoint\"]).eval().to(device)\n    dict_features_cuda = {}\n    dict_shapes = {}\n    dict_images = {}\n    #dict_fname_shapes = {}\n    for X in dataloader:\n        image, idx, ori_shape = X\n        image = image[0].to(device)\n        fname = img_fnames[idx]\n        #dict_fname_shapes[fname] = ori_shape\n        key = fname.split('/')[-1]\n\n        with torch.no_grad(), torch.autocast(device_type=\"cuda\", dtype=torch.float16):\n            pred = superpoint({'image': image})\n            dict_features_cuda[key] = pred\n            dict_shapes[key] = ori_shape\n            dict_images[key] = image.half()\n    del superpoint\n    gc.collect()\n\n    #####################################################\n    # Matching keypoints\n    #####################################################\n    superglue = SuperGlue(sg_config[\"superglue\"]).eval().to(device)\n    weights = sg_config[\"superglue\"][\"weights\"]\n    cnt_pairs = 0\n\n    mkpt = {}\n    with h5py.File(file_keypoints, mode='w') as f_match:\n        for idx, (fname1, fname2) in enumerate(zip(fnames1, fnames2)):\n            key1, key2 = fname1.split('/')[-1], fname2.split('/')[-1]\n\n            data = {\"image0\": dict_images[key1], \"image1\": dict_images[key2]}\n            data = {**data, **{k+'0': v for k, v in dict_features_cuda[key1].items()}}\n            data = {**data, **{k+'1': v for k, v in dict_features_cuda[key2].items()}}\n            for k in data:\n                if isinstance(data[k], (list, tuple)):\n                    data[k] = torch.stack(data[k])\n            with torch.no_grad(), torch.autocast(device_type=\"cuda\", dtype=torch.float16):\n                pred = {**data, **superglue(data)}\n                pred = {k: v[0].detach().cpu().numpy().copy() for k, v in pred.items()}\n            mkpts1, mkpts2 = pred[\"keypoints0\"], pred[\"keypoints1\"]\n            matches, conf = pred[\"matches0\"], pred[\"matching_scores0\"]\n            valid = matches > -1\n            mkpts1 = mkpts1[valid]\n            mkpts2 = mkpts2[matches[valid]]\n            mconf = conf[valid]\n            ori_shape_1 = dict_shapes[key1][0].numpy()\n            ori_shape_2 = dict_shapes[key2][0].numpy()\n            # Scaling coords\n            mkpts1[:,0] = mkpts1[:,0] * ori_shape_1[1] / dict_images[key1].shape[3]   # X\n            mkpts1[:,1] = mkpts1[:,1] * ori_shape_1[0] / dict_images[key1].shape[2]   # Y\n            mkpts2[:,0] = mkpts2[:,0] * ori_shape_2[1] / dict_images[key2].shape[3]   # X\n            mkpts2[:,1] = mkpts2[:,1] * ori_shape_2[0] / dict_images[key2].shape[2]   # Y  \n            n_matches = mconf.shape[0]\n            mkpt[fname1] = np.concatenate([mkpt[fname1], mkpts1], axis=0).astype(np.float32) if fname1 in mkpt else mkpts1\n            mkpt[fname2] = np.concatenate([mkpt[fname2], mkpts2], axis=0).astype(np.float32) if fname2 in mkpt else mkpts2\n            group  = f_match.require_group(key1)\n            if n_matches >= min_matches:\n                group.create_dataset(key2, data=np.concatenate([mkpts1, mkpts2], axis=1).astype(np.float32))\n                cnt_pairs+=1\n                print (f'{key1}-{key2}: {n_matches} matches @ {cnt_pairs}th pair(superglue/{resize_to}/{weights})')            \n            else:\n                print (f'{key1}-{key2}: {n_matches} matches --> skipped')\n\n\n\n    gc.collect()\n    del superglue\n    del dict_features_cuda\n    del dict_images\n    torch.cuda.empty_cache()\n    gc.collect()\n    t=time() -t \n    print(f'Features matched in  {t:.4f} sec')\n    return t","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for imgIdx in range(maxImageCount):\n    fig, axs = plt.subplots(2, 2, figsize=(12, 12), gridspec_kw={'wspace':0.1, 'hspace':0.1}, squeeze=True)\n    axs = axs.ravel()\n    for ax in axs:\n        ax.set_axis_off()\n    \n    for resolutionIdx, resolution in enumerate(resolutionList):\n        displayImage = drawKeypoints(keypointsPerResolution[resolution][imgIdx], images[imgIdx].copy())\n        axs[resolutionIdx].imshow(displayImage)\n        axs[resolutionIdx].set_title(f\"SP: {resolution}. KPs: {len(keypointsPerResolution[resolution][imgIdx])}.\")\n    \n    displayImage = drawKeypoints(featuresXYSIFT[imgIdx], images[imgIdx].copy())\n    axs[-1].imshow(displayImage)\n    axs[-1].set_title(f\"SIFT. KPs: {len(featuresXYSIFT[imgIdx])}.\")","metadata":{"execution":{"iopub.status.busy":"2024-02-27T03:17:01.336816Z","iopub.execute_input":"2024-02-27T03:17:01.337282Z","iopub.status.idle":"2024-02-27T03:17:26.491959Z","shell.execute_reply.started":"2024-02-27T03:17:01.337243Z","shell.execute_reply":"2024-02-27T03:17:26.491079Z"},"trusted":true},"execution_count":null,"outputs":[]}]}