{"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":"none","dataSources":[{"sourceType":"competition","sourceId":71885,"databundleVersionId":8143495}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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\nimport 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 = \"cuda\" if torch.cuda.is_available() else \"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 = np.array([qx, qy, qz, qw])\n        quat = quat / np.linalg.norm(quat)\n\n        pose = np.concatenate([quat, [tx, ty, tz]])\n        pose_dict[base] = pose\n\n    data = []\n    for img_name in os.listdir(img_dir):\n        base = os.path.splitext(img_name)[0]\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# ==============================\n# BUILD DATASET\n# ==============================\ndata = []\nfor s in os.listdir(ROOT):\n    data.extend(parse_scene(s))\n\nprint(\"TOTAL IMAGES:\", len(data))\n\n# Use more data\ndata = data[:1000]\n\nnp.random.shuffle(data)\n\n# ==============================\n# AUGMENTATIONS\n# ==============================\ntrain_transform = A.Compose([\n    A.Resize(224,224),\n    A.HorizontalFlip(p=0.5),\n    A.RandomBrightnessContrast(p=0.3),\n    A.Rotate(limit=20, p=0.3),\n    A.Normalize(),\n    ToTensorV2()\n])\n\nval_transform = A.Compose([\n    A.Resize(224,224),\n    A.Normalize(),\n    ToTensorV2()\n])\n\n# ==============================\n# DATASET CLASS\n# ==============================\nclass PoseDS(Dataset):\n    def __init__(self, data, transform):\n        self.data = data\n        self.transform = transform\n\n    def __len__(self):\n        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\n        img = self.transform(image=img)['image']\n\n        return img, torch.tensor(target, dtype=torch.float32)\n\n# ==============================\n# SPLIT\n# ==============================\nsplit = int(0.8 * len(data))\n\ntrain_loader = DataLoader(\n    PoseDS(data[:split], train_transform),\n    batch_size=32,\n    shuffle=True\n)\n\nval_loader = DataLoader(\n    PoseDS(data[split:], val_transform),\n    batch_size=32\n)\n\n# ==============================\n# MODEL\n# ==============================\nclass PoseNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = timm.create_model(\n            \"efficientnet_b0\",\n            pretrained=True,\n            num_classes=0\n        )\n        self.head = nn.Sequential(\n            nn.Linear(1280, 512),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, 7)\n        )\n\n    def forward(self, x):\n        x = self.backbone(x)\n        x = self.head(x)\n        return x\n\nmodel = PoseNet().to(DEVICE)\n\n# ==============================\n# LOSS FUNCTION\n# ==============================\ndef pose_loss(pred, target):\n    pq, pt = pred[:, :4], pred[:, 4:]\n    tq, tt = target[:, :4], target[:, 4:]\n\n    pq = pq / torch.norm(pq, dim=1, keepdim=True)\n    tq = tq / torch.norm(tq, dim=1, keepdim=True)\n\n    rot_loss = torch.mean(1 - torch.sum(pq * tq, dim=1) ** 2)\n    trans_loss = nn.functional.mse_loss(pt, tt)\n\n    return rot_loss + 10.0 * trans_loss\n\n\n# ==============================\n# OPTIMIZER\n# ==============================\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='min', patience=3, factor=0.5\n)\n\n# ==============================\n# TRAINING\n# ==============================\nEPOCHS = 40\ntrain_loss, val_loss, accs = [], [], []\n\nfor epoch in range(EPOCHS):\n    model.train()\n    tl = 0\n\n    for x, y in tqdm(train_loader):\n        x, y = x.to(DEVICE), y.to(DEVICE)\n\n        optimizer.zero_grad()\n        out = model(x)\n        loss = pose_loss(out, y)\n        loss.backward()\n        optimizer.step()\n\n        tl += loss.item()\n\n    tl /= len(train_loader)\n\n    # Validation\n    model.eval()\n    vl = 0\n    err = []\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            loss = pose_loss(out, y)\n            vl += loss.item()\n\n            mae = torch.mean(torch.abs(out - y)).item()\n            err.append(mae)\n\n    vl /= len(val_loader)\n    scheduler.step(vl)\n\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 {epoch+1}: Train={tl:.4f} Val={vl:.4f} Acc={acc:.4f}\")\n\n# ==============================\n# PLOT LOSS\n# ==============================\nplt.plot(train_loss, label=\"Train\")\nplt.plot(val_loss, label=\"Val\")\nplt.legend()\nplt.title(\"Loss Curve\")\nplt.show()\n\n# ==============================\n# PLOT ACCURACY\n# ==============================\nplt.plot(accs)\nplt.title(\"Validation Accuracy\")\nplt.show()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-09T06:30:04.942777Z","iopub.execute_input":"2026-04-09T06:30:04.943173Z"}},"outputs":[],"execution_count":null}]}