{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport os\nimport cv2\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\nimport torchvision\nfrom torchvision import transforms ","metadata":{"execution":{"iopub.status.busy":"2023-05-14T21:48:35.100267Z","iopub.execute_input":"2023-05-14T21:48:35.101425Z","iopub.status.idle":"2023-05-14T21:48:35.111768Z","shell.execute_reply.started":"2023-05-14T21:48:35.101345Z","shell.execute_reply":"2023-05-14T21:48:35.110409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"root = '/kaggle/input/image-matching-challenge-2023'\ntrain_label_file = 'train_labels.csv'\ntrain_path = '/kaggle/input/image-matching-challenge-2023/train/'\ntest_path ='/kaggle/input/image-matching-challenge-2023/test/'","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:07:37.430881Z","iopub.execute_input":"2023-05-14T22:07:37.431321Z","iopub.status.idle":"2023-05-14T22:07:37.437497Z","shell.execute_reply.started":"2023-05-14T22:07:37.431281Z","shell.execute_reply":"2023-05-14T22:07:37.436125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_datasets(root,path):\n    file_path = os.path.join(root,path)\n    df = pd.read_csv(file_path)\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:07:40.557282Z","iopub.execute_input":"2023-05-14T22:07:40.558388Z","iopub.status.idle":"2023-05-14T22:07:40.564335Z","shell.execute_reply.started":"2023-05-14T22:07:40.558325Z","shell.execute_reply":"2023-05-14T22:07:40.562988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_datasets(train_path,train_label_file).head()","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:07:41.266352Z","iopub.execute_input":"2023-05-14T22:07:41.266763Z","iopub.status.idle":"2023-05-14T22:07:41.287758Z","shell.execute_reply.started":"2023-05-14T22:07:41.266729Z","shell.execute_reply":"2023-05-14T22:07:41.286788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = A.Compose([\n    A.Resize(224,224),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), max_pixel_value=255.0, always_apply=False, p=1.0),\n    ToTensorV2()\n   ])\n\nvalidation_transform = A.Compose([\n    A.Resize(224,224),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225), max_pixel_value=255.0, always_apply=False, p=1.0),\n    ToTensorV2()\n    ])","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:40:56.989271Z","iopub.execute_input":"2023-05-14T22:40:56.990026Z","iopub.status.idle":"2023-05-14T22:40:56.999779Z","shell.execute_reply.started":"2023-05-14T22:40:56.989984Z","shell.execute_reply":"2023-05-14T22:40:56.998632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImageDataset(Dataset):\n    def __init__(self, root, img_path, rotation_label, translation_label, transforms=None):\n        self.root = root\n        self.img_path = img_path\n        self.rotation_label = rotation_label\n        self.translation_label = translation_label\n        self.transforms = transforms\n        \n    def __getitem__(self, index):\n        img_path = self.img_path[index]\n        img_path = self.root + img_path\n        image = cv2.imread(img_path)\n        \n        if self.transforms is not None:\n            image = self.transforms(image=image)['image']\n          \n        rotation_label = self.rotation_label[index].split(\";\")\n        rotation_label = list(map(float, rotation_label))\n        \n        translation_label = self.translation_label[index].split(\";\")\n        translation_label = list(map(float, translation_label))\n        \n        return  image, np.array(rotation_label), np.array(translation_label)\n        \n    def __len__(self):\n        return len(self.img_path)","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:42:44.962349Z","iopub.execute_input":"2023-05-14T22:42:44.963099Z","iopub.status.idle":"2023-05-14T22:42:44.975833Z","shell.execute_reply.started":"2023-05-14T22:42:44.963059Z","shell.execute_reply":"2023-05-14T22:42:44.974343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_validation_set(df):\n    validtion_data = df.sample(frac=0.3)\n    training_data =  df[~df[\"image_path\"].isin(validtion_data[\"image_path\"])]\n    return training_data,validtion_data\n\n\ntraining_data,validation_data = get_train_validation_set(get_datasets(train_path,train_label_file))","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:42:45.493586Z","iopub.execute_input":"2023-05-14T22:42:45.494624Z","iopub.status.idle":"2023-05-14T22:42:45.516069Z","shell.execute_reply.started":"2023-05-14T22:42:45.494580Z","shell.execute_reply":"2023-05-14T22:42:45.515038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_translation_rotation(df):\n    df[\"rotation_matrix_split\"] = df.apply(lambda x:list(map(float, x[\"rotation_matrix\"].split(\";\"))), axis=1)\n    df[\"translation_vector_split\"] = df.apply(lambda x:list(map(float, x[\"translation_vector\"].split(\";\"))), axis=1)\n    rotation_value = np.array(df[\"rotation_matrix_split\"].tolist())\n    translation_value = np.array(df[\"translation_vector_split\"].tolist())\n    \n    return translation_value,rotation_value\n\ntranslation_value,rotation_value = get_translation_rotation(get_datasets(train_path,train_label_file))","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:42:45.994058Z","iopub.execute_input":"2023-05-14T22:42:45.994465Z","iopub.status.idle":"2023-05-14T22:42:46.024994Z","shell.execute_reply.started":"2023-05-14T22:42:45.994429Z","shell.execute_reply":"2023-05-14T22:42:46.023968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_dataset(train_path,training_data):\n    train_dataset = ImageDataset(train_path, \n                        training_data[\"image_path\"].tolist(), \n                        training_data[\"rotation_matrix\"].tolist(), \n                        training_data[\"translation_vector\"].tolist(), \n                        transforms=train_transform)\n    return train_dataset","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:42:46.497888Z","iopub.execute_input":"2023-05-14T22:42:46.498278Z","iopub.status.idle":"2023-05-14T22:42:46.504813Z","shell.execute_reply.started":"2023-05-14T22:42:46.498243Z","shell.execute_reply":"2023-05-14T22:42:46.503623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_validation_dataset(train_path,validation_data):    \n    validation_dataset = ImageDataset(train_path, \n                        validation_data[\"image_path\"].tolist(), \n                        validation_data[\"rotation_matrix\"].tolist(), \n                        validation_data[\"translation_vector\"].tolist(), \n                        transforms=validation_transform)\n    return validation_dataset","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:42:47.023750Z","iopub.execute_input":"2023-05-14T22:42:47.024631Z","iopub.status.idle":"2023-05-14T22:42:47.031026Z","shell.execute_reply.started":"2023-05-14T22:42:47.024581Z","shell.execute_reply":"2023-05-14T22:42:47.029735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_dataloader(train_dataset):\n    train_loader = DataLoader(train_dataset, batch_size = 4, shuffle=True, num_workers=2)\n    return train_loader\n    \ndef get_validation_dataloader(validation_dataset):\n    validation_loader = DataLoader(validation_dataset, batch_size=4, shuffle=False)\n    return validation_loader","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:42:47.717488Z","iopub.execute_input":"2023-05-14T22:42:47.717888Z","iopub.status.idle":"2023-05-14T22:42:47.724772Z","shell.execute_reply.started":"2023-05-14T22:42:47.717853Z","shell.execute_reply":"2023-05-14T22:42:47.723404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader=get_train_dataloader(get_train_dataset(train_path,training_data))\nvalidation_loader = get_validation_dataloader(get_validation_dataset(train_path,validation_data))","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:42:48.176889Z","iopub.execute_input":"2023-05-14T22:42:48.177933Z","iopub.status.idle":"2023-05-14T22:42:48.185815Z","shell.execute_reply.started":"2023-05-14T22:42:48.177874Z","shell.execute_reply":"2023-05-14T22:42:48.184528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImagemtachingModel(torch.nn.Module):\n    def __init__(self, dropout=0.2):\n        super(ImagemtachingModel, self).__init__()\n        self.layer_one = nn.Sequential(\n            nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(64),\n            nn.ReLU(),\n            nn.Conv2d(64, 64, kernel_size=3, stride=1, padding=1), \n            nn.BatchNorm2d(64), \n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.MaxPool2d(kernel_size=2))\n        \n        self.layer_two = nn.Sequential(\n            nn.Conv2d(64, 128, kernel_size=3, stride=1, padding=1), \n            nn.BatchNorm2d(128),\n            nn.ReLU(),\n            nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(128), \n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.MaxPool2d(kernel_size=2))\n        \n        self.layer_three = nn.Sequential(\n            nn.Conv2d(128, 256, kernel_size=3, stride=1, padding=1), \n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(256),\n            nn.ReLU(),\n            nn.Conv2d(256, 256, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(256), \n            nn.ReLU(), \n            nn.Dropout(dropout),\n            nn.MaxPool2d(kernel_size=2))\n        \n        self.layer_four = nn.Sequential(\n            nn.Conv2d(256, 512, kernel_size=3, stride=1, padding=1), \n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(512), \n            nn.ReLU(), \n            nn.Dropout(dropout),\n            nn.MaxPool2d(kernel_size=2))\n        \n        self.layer_five = nn.Sequential(\n            nn.Conv2d(512, 512, kernel_size=3, stride=1, padding=1), \n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(512),\n            nn.ReLU(),\n            nn.Conv2d(512, 512, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(512), \n            nn.ReLU(), \n            nn.Dropout(dropout),\n            nn.MaxPool2d(kernel_size=2), \n            nn.AvgPool2d(kernel_size=7))\n        \n        self.layer_six = nn.Sequential(\n            nn.Linear(512, 512),\n            nn.ReLU(),\n        )\n        self.rotation_out = nn.Linear(512, 9)\n        self.tanh = torch.nn.Tanh()\n        \n        self.translation_out = torch.nn.Linear(512, 3)\n        \n    def forward(self, x):\n        x = self.layer_one(x)\n        x = self.layer_two(x)\n        x = self.layer_three(x)\n        x = self.layer_four(x)\n        x = self.layer_five(x)\n        \n        x = x.view(-1, 512)\n        x = self.layer_six(x)\n        \n        rotation_out = self.rotation_out(x)\n        rotation_out = self.tanh(rotation_out)\n        \n        translation_out = self.translation_out(x)\n        \n        return rotation_out, translation_out","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:43:18.244921Z","iopub.execute_input":"2023-05-14T22:43:18.245313Z","iopub.status.idle":"2023-05-14T22:43:18.272557Z","shell.execute_reply.started":"2023-05-14T22:43:18.245278Z","shell.execute_reply":"2023-05-14T22:43:18.271067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:43:18.718799Z","iopub.execute_input":"2023-05-14T22:43:18.719183Z","iopub.status.idle":"2023-05-14T22:43:18.726018Z","shell.execute_reply.started":"2023-05-14T22:43:18.719150Z","shell.execute_reply":"2023-05-14T22:43:18.724689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ImagemtachingModel(dropout=0.2)\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:43:19.190928Z","iopub.execute_input":"2023-05-14T22:43:19.191778Z","iopub.status.idle":"2023-05-14T22:43:19.367243Z","shell.execute_reply.started":"2023-05-14T22:43:19.191740Z","shell.execute_reply":"2023-05-14T22:43:19.366060Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"l1_distance = torch.nn.L1Loss()\noptimizer = torch.optim.Adam(params = model.parameters(), lr = 0.001)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=2, threshold_mode='abs', min_lr=1e-8, verbose=True)\nbest_loss = 1000000000\nepochs = 10\nbest_model = None","metadata":{"execution":{"iopub.status.busy":"2023-05-14T22:43:19.606631Z","iopub.execute_input":"2023-05-14T22:43:19.606973Z","iopub.status.idle":"2023-05-14T22:43:19.614344Z","shell.execute_reply.started":"2023-05-14T22:43:19.606942Z","shell.execute_reply":"2023-05-14T22:43:19.613182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for epoch in range(1, epochs+1):\n    \n    train_loss = []\n    rot_loss = []\n    trans_loss = []\n    \n    val_loss = []\n    val_rot_loss = []\n    val_trans_loss = []\n    \n    for imgs, rotation_labels, translation_labels in tqdm(train_loader):\n        model.train()\n        optimizer.zero_grad()\n        \n        imgs = imgs.to(device)\n        rotation_labels = rotation_labels.to(device)\n        translation_labels = translation_labels.to(device)\n        \n        rotation_output, translation_output = model(imgs)\n        \n        rotation_loss = l1_distance(rotation_output, rotation_labels)\n        translation_loss = l1_distance(translation_output, translation_labels)\n        loss = rotation_loss + translation_loss\n        loss.backward()\n        \n        optimizer.step()\n        \n        train_loss.append(loss.item())\n        rot_loss.append(rotation_loss.item())\n        trans_loss.append(translation_loss.item())\n\n    for imgs, rotation_labels, translation_labels in tqdm(validation_loader):\n        model.eval()\n        \n        imgs = imgs.to(device)\n        rotation_labels = rotation_labels.to(device)\n        translation_labels = translation_labels.to(device)\n        \n        rotation_output, translation_output = model(imgs)\n        rotation_loss = l1_distance(rotation_output, rotation_labels)\n        translation_loss = l1_distance(translation_output, translation_labels)\n        loss = rotation_loss + translation_loss\n        \n        val_loss.append(loss.item())\n        val_rot_loss.append(rotation_loss.item())\n        val_trans_loss.append(translation_loss.item())\n        \n    mtrain_loss = np.mean(train_loss)\n    mval_loss = np.mean(val_loss)\n    mtrain_rot_loss = np.mean(rot_loss)\n    mtrain_trans_loss = np.mean(trans_loss)\n    mval_rot_loss = np.mean(val_rot_loss)\n    mval_trans_loss = np.mean(val_trans_loss)\n    \n\n\n    if scheduler is not None:\n        scheduler.step(mval_loss)\n\n    if best_loss < mval_loss:\n        best_loss = mval_loss\n        best_model = model\n        \n    print(f'Epoch [{epoch}], Train Loss : [{mtrain_loss:.5f}] \\\n    Train Rotation Loss : [{mtrain_rot_loss:.5f}] Train Translation Loss : [{mtrain_trans_loss:.5f}] \\\n          Val Loss : [{mval_loss:.5f}] Val Rotation Loss : [{mval_rot_loss:.5f}] Val Translation Loss : [{mval_trans_loss:.5f}]')","metadata":{"execution":{"iopub.status.busy":"2023-05-14T23:02:56.893033Z","iopub.execute_input":"2023-05-14T23:02:56.893616Z","iopub.status.idle":"2023-05-14T23:09:57.083419Z","shell.execute_reply.started":"2023-05-14T23:02:56.893577Z","shell.execute_reply":"2023-05-14T23:09:57.082053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def arr_to_str(a):\n    return ';'.join([str(x) for x in a.reshape(-1)])","metadata":{"execution":{"iopub.status.busy":"2023-05-14T23:10:02.229837Z","iopub.execute_input":"2023-05-14T23:10:02.231017Z","iopub.status.idle":"2023-05-14T23:10:02.237846Z","shell.execute_reply.started":"2023-05-14T23:10:02.230962Z","shell.execute_reply":"2023-05-14T23:10:02.236348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv(\"/kaggle/input/image-matching-challenge-2023/sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-05-14T23:10:02.737315Z","iopub.execute_input":"2023-05-14T23:10:02.738549Z","iopub.status.idle":"2023-05-14T23:10:02.749666Z","shell.execute_reply.started":"2023-05-14T23:10:02.738501Z","shell.execute_reply":"2023-05-14T23:10:02.748410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_rotation_list = []\ntest_translation_list = []\ntest_img_list = []\ntest_dataset_list = []\ntest_scene_list = []\nfor i in range(len(submission)):\n    test_row = submission.iloc[i]\n    test_img = test_row[\"image_path\"]\n    test_dataset = test_row[\"dataset\"]\n    test_scene = test_row[\"scene\"]\n    img_path = f\"{root}/test/{test_img}\"\n    \n    try:\n        image = cv2.imread(img_path)\n        image = test_transforms(image=image)['image']\n\n        rotation, translation = model(image)\n\n        rotation = arr_to_str(rotation.detach().cpu().numpy())\n        translation = arr_to_str(translation.detach().cpu().numpy())\n    except:\n        rotation = \"1.0;0.0;0.0;0.0;1.0;0.0;0.0;0.0;1.0\"\n        translation = \"0.0;0.0;0.0\"\n    \n    test_rotation_list.append(rotation)\n    test_translation_list.append(translation)\n    test_img_list.append(test_img)\n    test_dataset_list.append(test_dataset)\n    test_scene_list.append(test_scene)","metadata":{"execution":{"iopub.status.busy":"2023-05-14T23:10:03.291898Z","iopub.execute_input":"2023-05-14T23:10:03.292878Z","iopub.status.idle":"2023-05-14T23:10:03.305273Z","shell.execute_reply.started":"2023-05-14T23:10:03.292823Z","shell.execute_reply":"2023-05-14T23:10:03.304057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"my_submission = pd.DataFrame()\nmy_submission[\"image_path\"] = test_img_list\nmy_submission[\"dataset\"] = test_dataset_list\nmy_submission[\"scene\"] = test_scene_list\nmy_submission[\"rotation_matrix\"] = test_rotation_list\nmy_submission[\"translation_vector\"] = test_translation_list\n\nmy_submission.to_csv(\"submission.csv\",index=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-14T23:10:04.868072Z","iopub.execute_input":"2023-05-14T23:10:04.868602Z","iopub.status.idle":"2023-05-14T23:10:04.881940Z","shell.execute_reply.started":"2023-05-14T23:10:04.868561Z","shell.execute_reply":"2023-05-14T23:10:04.880545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv(\"/kaggle/working/submission.csv\")\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-05-14T23:10:29.194094Z","iopub.execute_input":"2023-05-14T23:10:29.194635Z","iopub.status.idle":"2023-05-14T23:10:29.217833Z","shell.execute_reply.started":"2023-05-14T23:10:29.194586Z","shell.execute_reply":"2023-05-14T23:10:29.216235Z"},"trusted":true},"execution_count":null,"outputs":[]}]}