{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":71885,"databundleVersionId":8143495}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ==============================\n# INSTALL\n# ==============================\n!pip install albumentations timm scipy -q\n\n# ==============================\n# IMPORTS\n# ==============================\nimport os, cv2, torch, timm\nimport numpy as np, matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom scipy.spatial.transform import Rotation as R\n\nDEVICE = \"cpu\"\n\nBASE = \"/kaggle/input/competitions/image-matching-challenge-2024\"\nROOT = f\"{BASE}/train\"\n\n# ==============================\n# PARSER\n# ==============================\ndef parse_scene(scene):\n    img_dir = os.path.join(ROOT, scene, \"images\")\n    sfm_file = os.path.join(ROOT, scene, \"sfm/images.txt\")\n\n    if not os.path.exists(img_dir) or not os.path.exists(sfm_file):\n        return []\n\n    pose_dict = {}\n\n    with open(sfm_file) as f:\n        lines = f.readlines()\n\n    for line in lines:\n        if line.startswith(\"#\"):\n            continue\n\n        parts = line.strip().split()\n        if len(parts) < 10:\n            continue\n\n        name = os.path.basename(parts[-1])\n        base = os.path.splitext(name)[0]\n\n        qw,qx,qy,qz = map(float, parts[1:5])\n        tx,ty,tz = map(float, parts[5:8])\n\n        quat = R.from_quat([qx,qy,qz,qw]).as_quat()\n        pose = np.concatenate([quat, [tx,ty,tz]])\n\n        pose_dict[base] = pose\n\n    data = []\n\n    for img_name in os.listdir(img_dir):\n        base = os.path.splitext(img_name)[0]\n\n        if base in pose_dict:\n            img_path = f\"{scene}/images/{img_name}\"\n            data.append((img_path, pose_dict[base]))\n\n    print(f\"{scene} → {len(data)} images\")\n    return data\n\n# ==============================\n# BUILD DATASET\n# ==============================\ndata = []\nfor s in os.listdir(ROOT):\n    data.extend(parse_scene(s))\n\nprint(\"TOTAL IMAGES:\", len(data))\ndata = data[:150]\n\n# ==============================\n# DATASET CLASS\n# ==============================\ntransform = A.Compose([\n    A.Resize(224, 224),\n    A.Normalize(),\n    ToTensorV2()\n])\n\nclass PoseDS(Dataset):\n    def __init__(self, data):\n        self.data = data\n\n    def __len__(self): return len(self.data)\n\n    def __getitem__(self, i):\n        path, target = self.data[i]\n\n        img = cv2.imread(os.path.join(ROOT, path))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = transform(image=img)['image']\n\n        return img, torch.tensor(target, dtype=torch.float32)\n\n# ==============================\n# SPLIT\n# ==============================\nnp.random.shuffle(data)\nsplit = int(0.8 * len(data))\n\ntrain_loader = DataLoader(PoseDS(data[:split]), batch_size=16, shuffle=True)\nval_loader   = DataLoader(PoseDS(data[split:]), batch_size=16)\n\n# ==============================\n# MODEL\n# ==============================\nmodel = timm.create_model(\"resnet18\", pretrained=True, num_classes=7).to(DEVICE)\n\nfeatures = []\ndef hook(m, i, o):\n    features.clear()\n    features.append(o.detach().cpu())\n\nmodel.conv1.register_forward_hook(hook)\n\n# ==============================\n# LOSS FUNCTION\n# ==============================\ndef loss_fn(p, t):\n    pq, pt = p[:, :4], p[:, 4:]\n    tq, tt = t[:, :4], t[:, 4:]\n\n    pq = pq / torch.norm(pq, dim=1, keepdim=True)\n\n    q_loss = torch.mean(1 - torch.sum(pq * tq, dim=1) ** 2)\n    t_loss = nn.functional.mse_loss(pt, tt)\n\n    return q_loss + t_loss\n\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n\n# ==============================\n# TRAINING\n# ==============================\nEPOCHS = 15\ntrain_loss, val_loss, accs = [], [], []\n\nfor e in range(EPOCHS):\n    model.train()\n    tl = 0\n\n    for x, y in train_loader:\n        x, y = x.to(DEVICE), y.to(DEVICE)\n\n        optimizer.zero_grad()\n        out = model(x)\n        loss = loss_fn(out, y)\n        loss.backward()\n        optimizer.step()\n\n        tl += loss.item()\n\n    tl /= len(train_loader)\n\n    model.eval()\n    vl, err = 0, []\n\n    with torch.no_grad():\n        for x, y in val_loader:\n            x, y = x.to(DEVICE), y.to(DEVICE)\n            out = model(x)\n\n            vl += loss_fn(out, y).item()\n            err.append(torch.mean(torch.abs(out - y)).item())\n\n    vl /= len(val_loader)\n    acc = 1 / (1 + np.mean(err))\n\n    train_loss.append(tl)\n    val_loss.append(vl)\n    accs.append(acc)\n\n    print(f\"Epoch {e+1} → Train:{tl:.4f} Val:{vl:.4f} Acc:{acc:.4f}\")\n\n# ==============================\n# LOSS GRAPH\n# ==============================\nplt.plot(train_loss, label=\"Train\")\nplt.plot(val_loss, label=\"Val\")\nplt.legend()\nplt.title(\"Loss Curve\")\nplt.show()\n\n# ==============================\n# FEATURE MAP\n# ==============================\nimg_path, _ = data[0]\nimg = cv2.imread(os.path.join(ROOT, img_path))\nimg = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\ninp = transform(image=img)['image'].unsqueeze(0).to(DEVICE)\n\nwith torch.no_grad():\n    _ = model(inp)\n\nfm = features[0][0]\n\nplt.figure(figsize=(10, 5))\nfor i in range(6):\n    plt.subplot(2, 3, i + 1)\n    plt.imshow(fm[i], cmap='gray')\n    plt.axis('off')\n\nplt.suptitle(\"Feature Maps (Edges)\")\nplt.show()\n\n# ==============================\n# SIFT KEYPOINT MATCHING\n# ==============================\nIMG1_PATH = \"/kaggle/input/competitions/image-matching-challenge-2024/test/church/images/00001.png\"\nIMG2_PATH = \"/kaggle/input/competitions/image-matching-challenge-2024/test/church/images/00018.png\"\n\nscene = data[0][0].split(\"/\")[0]\nscene_imgs = [p for p, _ in data if p.startswith(scene)]\n\nimg1 = cv2.imread(IMG1_PATH)\nimg2 = cv2.imread(IMG2_PATH)\n\ng1 = cv2.cvtColor(img1, cv2.COLOR_BGR2GRAY)\ng2 = cv2.cvtColor(img2, cv2.COLOR_BGR2GRAY)\n\n\nsift = cv2.SIFT_create(nfeatures=2000, contrastThreshold=0.04, edgeThreshold=10)\n\nkp1, des1 = sift.detectAndCompute(g1, None)\nkp2, des2 = sift.detectAndCompute(g2, None)\n\nprint(f\"SIFT keypoints: img1={len(kp1)}, img2={len(kp2)}\")\n\n\nFLANN_INDEX_KDTREE = 1\nindex_params  = dict(algorithm=FLANN_INDEX_KDTREE, trees=5)\nsearch_params = dict(checks=50)  \n\nflann = cv2.FlannBasedMatcher(index_params, search_params)\nmatches = flann.knnMatch(des1, des2, k=2)\n\n# Lowe's ratio test — standard threshold 0.75\ngood = [m for m, n in matches if m.distance < 0.75 * n.distance]\nprint(f\"Good matches after Lowe's ratio test: {len(good)}\")\n\npts1 = np.float32([kp1[m.queryIdx].pt for m in good])\npts2 = np.float32([kp2[m.trainIdx].pt for m in good])\n\nh1, w1 = g1.shape\nfocal  = max(w1, h1)           # rough approximation\npp     = (w1 / 2.0, h1 / 2.0)\n\nE, mask = cv2.findEssentialMat(\n    pts1, pts2,\n    focal=focal, pp=pp,\n    method=cv2.RANSAC,\n    prob=0.999,\n    threshold=1.0\n)\n\ninliers = int(mask.sum())\nprint(f\"RANSAC inliers: {inliers}/{len(good)} ({100*inliers/max(len(good),1):.1f}%)\")\n\n# Filter to inlier matches only\ngood_inliers = [g for g, m in zip(good, mask.ravel()) if m]\n\n# Optional: recover relative pose (R, t) from E\n_, rot, tvec, _ = cv2.recoverPose(E, pts1, pts2, focal=focal, pp=pp)\nprint(f\"Recovered rotation (Rodrigues): {cv2.Rodrigues(rot)[0].ravel()}\")\nprint(f\"Recovered translation (unit):   {tvec.ravel()}\")\n\n# --- Draw top-40 inlier matches ---\ndraw_params = dict(\n    matchColor=(0, 255, 0),        # green for inliers\n    singlePointColor=(255, 0, 0),  # red for unmatched\n    matchesMask=None,\n    flags=cv2.DrawMatchesFlags_NOT_DRAW_SINGLE_POINTS\n)\n\nmatch_img = cv2.drawMatches(img1, kp1, img2, kp2, good_inliers[:40], None, **draw_params)\n\nplt.figure(figsize=(14, 6))\nplt.imshow(cv2.cvtColor(match_img, cv2.COLOR_BGR2RGB))\nplt.title(f\"SIFT Keypoint Matching + RANSAC  |  {len(good_inliers)} inliers shown (top 40 drawn)\")\nplt.axis(\"off\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-15T04:27:20.347048Z","iopub.execute_input":"2026-04-15T04:27:20.347350Z","iopub.status.idle":"2026-04-15T04:30:48.089021Z","shell.execute_reply.started":"2026-04-15T04:27:20.347321Z","shell.execute_reply":"2026-04-15T04:30:48.088077Z"}},"outputs":[],"execution_count":null}]}