{"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":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        #print(os.path.join(dirname, filename))\n        pass\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-12T15:56:32.17221Z","iopub.execute_input":"2023-06-12T15:56:32.172768Z","iopub.status.idle":"2023-06-12T15:56:36.167529Z","shell.execute_reply.started":"2023-06-12T15:56:32.172697Z","shell.execute_reply":"2023-06-12T15:56:36.166686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"FOLDER_NAME='SIFT Improv'  # Folder to save the checkpoints of the model during training Default: Model0_CNN\nUSE_CHECKPOINT=False      # Load from a checkpoint to run  Default: False\nCHECKPOINT_FILE=''        # File .pth to load checkpoint from   Default: ''\nEPOCHS=5             # Total number of epochs  Default: 100\nSTART_EPOCH=0             # Which epoch to start from   Default:0\nBATCH_SIZE=5\nSEED=33\n\n","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:36.168743Z","iopub.execute_input":"2023-06-12T15:56:36.169026Z","iopub.status.idle":"2023-06-12T15:56:36.173658Z","shell.execute_reply.started":"2023-06-12T15:56:36.169001Z","shell.execute_reply":"2023-06-12T15:56:36.172983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nfrom pathlib import Path\nimport wandb\nimport os\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\nimport torchvision\nfrom torchinfo import summary\nimport cv2\nimport matplotlib.pyplot as plt\ntorch.manual_seed(SEED)","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:36.174712Z","iopub.execute_input":"2023-06-12T15:56:36.175069Z","iopub.status.idle":"2023-06-12T15:56:40.242058Z","shell.execute_reply.started":"2023-06-12T15:56:36.175034Z","shell.execute_reply":"2023-06-12T15:56:40.240923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.environ['WANDB_API_KEY']='Your wandb key'\n\nWANDB_PROJECT='IMC2023 Translation'","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.244386Z","iopub.execute_input":"2023-06-12T15:56:40.245379Z","iopub.status.idle":"2023-06-12T15:56:40.249531Z","shell.execute_reply.started":"2023-06-12T15:56:40.245347Z","shell.execute_reply":"2023-06-12T15:56:40.248588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"main_path=Path('/kaggle/input/image-matching-challenge-2023')\ntrain_path= main_path / 'train'\ntrain_labels_path= train_path / 'train_labels.csv'\ntrain_labels=pd.read_csv(train_labels_path)\ntrain_labels.head(4)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.251266Z","iopub.execute_input":"2023-06-12T15:56:40.252158Z","iopub.status.idle":"2023-06-12T15:56:40.307704Z","shell.execute_reply.started":"2023-06-12T15:56:40.252117Z","shell.execute_reply":"2023-06-12T15:56:40.306743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"out_path=Path('/kaggle/working/')\nout_path=out_path / FOLDER_NAME\nprint(str(out_path))\n\nprint(os.path.isdir(out_path))\n\nif not (os.path.isdir(out_path)):\n    os.makedirs(str(out_path))","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.309013Z","iopub.execute_input":"2023-06-12T15:56:40.309332Z","iopub.status.idle":"2023-06-12T15:56:40.316013Z","shell.execute_reply.started":"2023-06-12T15:56:40.309305Z","shell.execute_reply":"2023-06-12T15:56:40.315036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_labels=train_labels\ndf_test_labels=train_labels","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.317404Z","iopub.execute_input":"2023-06-12T15:56:40.317775Z","iopub.status.idle":"2023-06-12T15:56:40.332433Z","shell.execute_reply.started":"2023-06-12T15:56:40.317739Z","shell.execute_reply":"2023-06-12T15:56:40.331436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#df_train_labels=train_labels.sample(frac=0.7, random_state=33)\n#df_test_labels=train_labels.drop(df_train_labels.index)\n\ndf_test_labels.head(5)\n\nlist_image_paths=df_train_labels['image_path'].tolist()\nlist_scenes=df_train_labels['scene'].tolist()\nlist_rotation_matrices=df_train_labels['rotation_matrix'].tolist()\nlist_transvectors=df_train_labels['translation_vector'].tolist()","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.333817Z","iopub.execute_input":"2023-06-12T15:56:40.334797Z","iopub.status.idle":"2023-06-12T15:56:40.344413Z","shell.execute_reply.started":"2023-06-12T15:56:40.334759Z","shell.execute_reply":"2023-06-12T15:56:40.343319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#print(list_image_paths)","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.345834Z","iopub.execute_input":"2023-06-12T15:56:40.34622Z","iopub.status.idle":"2023-06-12T15:56:40.355999Z","shell.execute_reply.started":"2023-06-12T15:56:40.346176Z","shell.execute_reply":"2023-06-12T15:56:40.355012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Savepoint","metadata":{}},{"cell_type":"code","source":"def save_checkpoint(EPOCH, model, optimizer, loss, PATH):\n    torch.save({\n            'epoch': EPOCH,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'loss': loss,\n            }, PATH)","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.360463Z","iopub.execute_input":"2023-06-12T15:56:40.36095Z","iopub.status.idle":"2023-06-12T15:56:40.367236Z","shell.execute_reply.started":"2023-06-12T15:56:40.360922Z","shell.execute_reply":"2023-06-12T15:56:40.366369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Creating image pair dataframe","metadata":{}},{"cell_type":"code","source":"list_image1=[]\nlist_image2=[]\nlist_R1=[]\nlist_R2=[]\nlist_T1=[]\nlist_T2=[]\ndf_train_pairs=pd.DataFrame()\n\nfor i in range(len(list_image_paths)):\n    #print('---------')\n    for j in range(len(list_image_paths)):\n        if(i !=j):\n            \n            image_path1=list_image_paths[i]\n\n            image_path2=list_image_paths[j]\n            R1=list_rotation_matrices[i]\n            R2=list_rotation_matrices[j]\n            T1=list_transvectors[i]\n            T2=list_transvectors[j]\n            scene1=list_scenes[i]\n            scene2=list_scenes[j]\n            \n            if scene1==scene2:\n                \n                list_image1.append(image_path1)\n                list_image2.append(image_path2)\n                list_R1.append(R1)\n                list_R2.append(R2)\n                list_T1.append(T1)\n                list_T2.append(T2)\n            else:\n                pass\n        \n            \n            \ndf_train_pairs['image_path1']=list_image1\ndf_train_pairs['image_path2']=list_image2\ndf_train_pairs['R1']=list_R1\ndf_train_pairs['R2']=list_R2\ndf_train_pairs['T1']=list_T1\ndf_train_pairs['T2']=list_T2\n\n\ndf_train_pairs_original=df_train_pairs","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.368632Z","iopub.execute_input":"2023-06-12T15:56:40.368919Z","iopub.status.idle":"2023-06-12T15:56:40.496341Z","shell.execute_reply.started":"2023-06-12T15:56:40.368895Z","shell.execute_reply":"2023-06-12T15:56:40.495174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"list_image_paths=df_test_labels['image_path'].tolist()\nlist_scenes=df_test_labels['scene'].tolist()\nlist_rotation_matrices=df_test_labels['rotation_matrix'].tolist()\nlist_transvectors=df_test_labels['translation_vector'].tolist()","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.497826Z","iopub.execute_input":"2023-06-12T15:56:40.498258Z","iopub.status.idle":"2023-06-12T15:56:40.504416Z","shell.execute_reply.started":"2023-06-12T15:56:40.49822Z","shell.execute_reply":"2023-06-12T15:56:40.503378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"list_image1=[]\nlist_image2=[]\nlist_R1=[]\nlist_R2=[]\nlist_T1=[]\nlist_T2=[]\ndf_test_pairs=pd.DataFrame()\n\nfor i in range(len(list_image_paths)):\n    #print('---------')\n    for j in range(len(list_image_paths)):\n        if(i !=j):\n            \n            image_path1=list_image_paths[i]\n\n            image_path2=list_image_paths[j]\n            R1=list_rotation_matrices[i]\n            R2=list_rotation_matrices[j]\n            T1=list_transvectors[i]\n            T2=list_transvectors[j]\n            scene1=list_scenes[i]\n            scene2=list_scenes[j]\n            \n            if scene1==scene2:\n                \n                list_image1.append(image_path1)\n                list_image2.append(image_path2)\n                list_R1.append(R1)\n                list_R2.append(R2)\n                list_T1.append(T1)\n                list_T2.append(T2)\n            else:\n                pass\n        \n            \n            \ndf_test_pairs['image_path1']=list_image1\ndf_test_pairs['image_path2']=list_image2\ndf_test_pairs['R1']=list_R1\ndf_test_pairs['R2']=list_R2\ndf_test_pairs['T1']=list_T1\ndf_test_pairs['T2']=list_T2\n\ndf_test_pairs_original=df_test_pairs","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.5058Z","iopub.execute_input":"2023-06-12T15:56:40.506222Z","iopub.status.idle":"2023-06-12T15:56:40.628264Z","shell.execute_reply.started":"2023-06-12T15:56:40.506185Z","shell.execute_reply":"2023-06-12T15:56:40.627239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make dataloaders","metadata":{}},{"cell_type":"code","source":"print(len(df_train_pairs))\nprint(len(df_test_pairs))","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.629745Z","iopub.execute_input":"2023-06-12T15:56:40.63024Z","iopub.status.idle":"2023-06-12T15:56:40.635488Z","shell.execute_reply.started":"2023-06-12T15:56:40.630204Z","shell.execute_reply":"2023-06-12T15:56:40.634598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class PairDataset(Dataset):\n    def __init__(self, \n                 df_pairs : pd.DataFrame,\n                 image_transform : transforms = None,\n                 feature_transform: transforms = None\n                ):\n        super().__init__()\n        self.list_image_path1=df_pairs['image_path1'].tolist()\n        self.list_image_path2=df_pairs['image_path2'].tolist()\n        \n        self.list_R1=df_pairs['R1'].tolist()\n        self.list_R2=df_pairs['R2'].tolist()\n        \n        self.list_T1=df_pairs['T1'].tolist()\n        self.list_T2=df_pairs['T2'].tolist()\n        \n        self.image_transform=image_transform\n        self.feature_transform=feature_transform\n        self.sift = cv2.SIFT_create()\n        self.orb = cv2.ORB_create()\n        self.matcher= cv2.BFMatcher(cv2.NORM_HAMMING, crossCheck=True)\n        \n    def __getitem__(self,index):\n        \n        image_path1='/kaggle/input/image-matching-challenge-2023/train/' + self.list_image_path1[index]\n        image_path2='/kaggle/input/image-matching-challenge-2023/train/' + self.list_image_path2[index]\n        \n        image1 = Image.open(image_path1)\n        image1 = self.image_transform(image1)\n        #image1=np.asarray(image1)\n        #train_keypoints1, train_descriptor1 = self.orb.detectAndCompute(image1, None)\n        \n        \n        \n        image2 = Image.open(image_path2)\n        image2= self.image_transform(image2)\n        \n        #image2=np.asarray(image2)\n        #train_keypoints2, train_descriptor2 = self.orb.detectAndCompute(image2, None)\n        \n        '''matches=self.matcher.match(train_descriptor1,train_descriptor2)\n        image1_points=[]\n        image2_points=[]\n\n        for match in matches:\n            image1_id=match.queryIdx\n            image2_id=match.trainIdx\n\n            image1_point=train_keypoints1[image1_id].pt\n            image2_point=train_keypoints2[image2_id].pt\n\n            image1_points.append(image1_point)\n            image2_points.append(image2_point)\n\n        image1_points= np.asarray(image1_points, dtype=np.float32)\n        image2_points= np.asarray(image2_points, dtype=np.float32)\n        \n        essential_matrix,mask = cv2.findEssentialMat(image1_points, image2_points)\n        _,rotmat, transvect,_= cv2.recoverPose(essential_matrix, image1_points, image2_points)'''\n        \n        rotation_matrix1=self.list_R1[index].split(';')\n        rotation_matrix1=np.array(rotation_matrix1, dtype=float)\n        rotation_matrix1=torch.from_numpy(rotation_matrix1).type(torch.float32)\n        \n        rotation_matrix2=self.list_R2[index].split(';')\n        rotation_matrix2=np.array(rotation_matrix2, dtype=float)\n        rotation_matrix2=torch.from_numpy(rotation_matrix2).type(torch.float32)\n        \n        \n        \n        \n        translation_vector1=self.list_T1[index].split(';')\n        translation_vector1=np.array(translation_vector1, dtype=float)\n        translation_vector1=torch.from_numpy(translation_vector1).type(torch.float32)\n        \n        translation_vector2=self.list_T2[index].split(';')\n        translation_vector2=np.array(translation_vector2, dtype=float)\n        translation_vector2=torch.from_numpy(translation_vector2).type(torch.float32)\n        \n        '''rotmat=torch.tensor(rotmat,dtype=torch.float32)\n        transvect=torch.tensor(transvect,dtype=torch.float32)\n        rotmat=rotmat.view(9)\n        transvect=transvect.view(3)\n        print(rotmat.shape)\n        print(transvect.shape)\n        print(rotation_matrix1.shape)\n        print(translation_vector1.shape)\n        print(rotation_matrix2.shape)\n        print(translation_vector2.shape)'''\n        return image1,image2, rotation_matrix1,translation_vector1,rotation_matrix2,translation_vector2\n        \n        \n    def __len__(self):\n        return len(self.list_image_path1)","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.637134Z","iopub.execute_input":"2023-06-12T15:56:40.637432Z","iopub.status.idle":"2023-06-12T15:56:40.65212Z","shell.execute_reply.started":"2023-06-12T15:56:40.637407Z","shell.execute_reply":"2023-06-12T15:56:40.651003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pretrained_vit_weights = torchvision.models.ViT_B_16_Weights.DEFAULT\nimage_transform = pretrained_vit_weights.transforms()\nfeature_transform=transforms.Compose([transforms.Resize(size=(128,128))])\ntrain_dataset=PairDataset(df_train_pairs, image_transform,feature_transform)\ntest_dataset=PairDataset(df_test_pairs, image_transform,feature_transform)","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.653462Z","iopub.execute_input":"2023-06-12T15:56:40.653792Z","iopub.status.idle":"2023-06-12T15:56:40.683562Z","shell.execute_reply.started":"2023-06-12T15:56:40.653764Z","shell.execute_reply":"2023-06-12T15:56:40.682411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Print a summary using torchinfo (uncomment for actual output)\n'''summary(model=model, \n         input_size=(1, 3, 224, 224), # (batch_size, color_channels, height, width)\n         # col_names=[\"input_size\"], # uncomment for smaller output\n         col_names=[\"input_size\", \"output_size\", \"num_params\", \"trainable\"],\n         col_width=20,\n         row_settings=[\"var_names\"]\n )'''","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.684845Z","iopub.execute_input":"2023-06-12T15:56:40.685196Z","iopub.status.idle":"2023-06-12T15:56:40.691154Z","shell.execute_reply.started":"2023-06-12T15:56:40.685153Z","shell.execute_reply":"2023-06-12T15:56:40.690231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Model tester with dataset\n\n'''image1,image2,rotmat1,transvect1,rotmat2,transvect2=train_dataset[0]\nimage1=image1.unsqueeze(0)\nimage2=image2.unsqueeze(0)\nprint(image1.shape)\nout=model(image1,image2)\nprint(out.shape)\nprint(out)'''","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.69375Z","iopub.execute_input":"2023-06-12T15:56:40.694231Z","iopub.status.idle":"2023-06-12T15:56:40.704977Z","shell.execute_reply.started":"2023-06-12T15:56:40.694203Z","shell.execute_reply":"2023-06-12T15:56:40.703994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader=DataLoader(dataset=train_dataset, batch_size=BATCH_SIZE, shuffle= True)\ntest_dataloader=DataLoader(dataset=test_dataset, batch_size=BATCH_SIZE, shuffle= False)","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.7062Z","iopub.execute_input":"2023-06-12T15:56:40.706757Z","iopub.status.idle":"2023-06-12T15:56:40.715342Z","shell.execute_reply.started":"2023-06-12T15:56:40.706729Z","shell.execute_reply":"2023-06-12T15:56:40.714661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### Model tester with dataloader\n\n'''CR,CT,R1,T1,R2,T2=next(iter(train_dataloader))\nprint(CR.shape)\nprint(CT.shape)\nprint(R1.shape)\nprint(T1.shape)\nprint(R2.shape)\nprint(T2.shape)\nprint('------------')\nR11=model(CR,CT)\nprint(R11.shape)'''\n","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.716291Z","iopub.execute_input":"2023-06-12T15:56:40.717005Z","iopub.status.idle":"2023-06-12T15:56:40.728028Z","shell.execute_reply.started":"2023-06-12T15:56:40.716949Z","shell.execute_reply":"2023-06-12T15:56:40.72711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make model","metadata":{}},{"cell_type":"code","source":"class DummyFCLayer(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n    def forward(self, x):\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.729364Z","iopub.execute_input":"2023-06-12T15:56:40.729827Z","iopub.status.idle":"2023-06-12T15:56:40.739338Z","shell.execute_reply.started":"2023-06-12T15:56:40.729792Z","shell.execute_reply":"2023-06-12T15:56:40.738551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SIFT_Translation_Regressor(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.vit_weights = torchvision.models.ViT_B_16_Weights.DEFAULT # requires torchvision >= 0.13, \"DEFAULT\" means best available\n\n\n        self.vit1 = torchvision.models.vit_b_16(weights=pretrained_vit_weights)\n        self.vit1.heads=nn.Linear(in_features=768,out_features=50)\n        \n        \n        for parameters in self.vit1.encoder.parameters():\n            parameters.requires_grad=False\n            \n        self.vit2 = torchvision.models.vit_b_16(weights=pretrained_vit_weights)\n        self.vit2.heads=nn.Linear(in_features=768,out_features=50)\n        \n        \n        for parameters in self.vit2.encoder.parameters():\n            parameters.requires_grad=False\n            \n        self.Regressor=nn.Linear(in_features=100,out_features=3)\n        \n\n\n    def forward(self,x,y):\n        feature_vector1=self.vit1(x)\n        feature_vector2=self.vit2(y)\n        feature_vector=torch.cat((feature_vector1,feature_vector2),dim=1)\n        rotmat=self.Regressor(feature_vector)\n       \n        return rotmat.squeeze()\n        \nmodel=SIFT_Translation_Regressor()","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:40.740084Z","iopub.execute_input":"2023-06-12T15:56:40.740349Z","iopub.status.idle":"2023-06-12T15:56:42.137994Z","shell.execute_reply.started":"2023-06-12T15:56:40.740326Z","shell.execute_reply":"2023-06-12T15:56:42.13562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#checkpoint = torch.load('/kaggle/input/imc2023-transvect-e4-s250/transvect_e4.pth')\n#model.load_state_dict(checkpoint['model_state_dict'])","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.139077Z","iopub.status.idle":"2023-06-12T15:56:42.139441Z","shell.execute_reply.started":"2023-06-12T15:56:42.139272Z","shell.execute_reply":"2023-06-12T15:56:42.139289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n\nloss_fn = nn.L1Loss()\n\noptimizer = torch.optim.Adam(params = model.parameters(), lr = 0.001)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=2, threshold_mode='abs', min_lr=1e-8, verbose=True)","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.140903Z","iopub.status.idle":"2023-06-12T15:56:42.141426Z","shell.execute_reply.started":"2023-06-12T15:56:42.141163Z","shell.execute_reply":"2023-06-12T15:56:42.141188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_mm(P,Q):\n    eps=1e-15\n    GT_SCALE = torch.linalg.norm(Q)\n    P = GT_SCALE * (P / (torch.linalg.norm(P) + eps))\n    err_t = min(torch.linalg.norm(Q - P), torch.linalg.norm(Q + P))\n    return err_t","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.142714Z","iopub.status.idle":"2023-06-12T15:56:42.143223Z","shell.execute_reply.started":"2023-06-12T15:56:42.142951Z","shell.execute_reply":"2023-06-12T15:56:42.142995Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Loop","metadata":{}},{"cell_type":"code","source":"def train_step(model: nn.Module,\n               dataloader: pd.DataFrame,\n               loss_fn: torch.nn.CrossEntropyLoss,\n               optimizer: torch.optim.SGD\n              ):\n    model.train()\n    train_trans_loss=0\n    image_transform=transforms.Compose([transforms.Resize(size=(224,224))])\n    feature_transform=transforms.Compose([transforms.Resize(size=(128,128))])\n    sift = cv2.SIFT_create()\n    \n    for batch, (image1,image2,rotation_matrix1,trans_vect1,rotation_matrix2,trans_vect2) in tqdm(enumerate(dataloader), total=len(dataloader), desc=\" Training\",position=0, leave=True):\n        \n\n            \n        \n        pred_transvect1=model(image1,image2)\n        \n        loss=get_mm(pred_transvect1,trans_vect1)\n        \n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n        trans_loss=loss.detach()\n        \n        train_trans_loss=train_trans_loss+trans_loss\n        \n    train_trans_loss =train_trans_loss / len(dataloader)\n    return  train_trans_loss\n        \n        \n        \n\n    \n    \n    ","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.144785Z","iopub.status.idle":"2023-06-12T15:56:42.145202Z","shell.execute_reply.started":"2023-06-12T15:56:42.144981Z","shell.execute_reply":"2023-06-12T15:56:42.145002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test Loop","metadata":{}},{"cell_type":"code","source":"def test_step(model: nn.Module,\n              dataloader: torch.utils.data.DataLoader,\n              loss_fn: torch.nn.CrossEntropyLoss\n             ):\n    \n    #model.eval()\n    test_trans_loss=0\n    \n    #with torch.inference_mode():\n    for batch, (image1,image2,rotation_matrix1,trans_vect1,rotation_matrix2,trans_vect2) in tqdm(enumerate(dataloader), total=len(dataloader), desc=\" Testing\",position=0, leave=True):\n\n        pred_transvect1=model(image1,image2)\n        loss_t=get_mm(pred_transvect1,trans_vect1)\n        loss=loss_t.detach()\n        test_trans_loss=test_trans_loss+loss\n            \n          \n    test_trans_loss =test_trans_loss / len(dataloader) \n    return test_trans_loss","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.146628Z","iopub.status.idle":"2023-06-12T15:56:42.147154Z","shell.execute_reply.started":"2023-06-12T15:56:42.146973Z","shell.execute_reply":"2023-06-12T15:56:42.146992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train epochs","metadata":{}},{"cell_type":"code","source":"def train(model: nn.Module,\n          train_dataloader: torch.utils.data.DataLoader,\n          test_dataloader: torch.utils.data.DataLoader,\n          loss_fn: torch.nn.CrossEntropyLoss,\n          optimizer: torch.optim.SGD,\n          epochs: int,\n          out_path: os.path\n         ):\n    results={\n             \"Train Translation Loss\" : [],\n             \"Test Translation Loss\" : []\n            }\n\n    for epoch in tqdm(range(START_EPOCH,EPOCHS), total=EPOCHS, desc='Epochs', position=0, leave=True):\n        print(epoch)\n        train_trans_loss= train_step(model=model,\n                                   dataloader=train_dataloader,\n                                   loss_fn=loss_fn,\n                                   optimizer=optimizer)\n        #print(train_rot_loss)\n        out_path_file=out_path / f\"epoch-{epoch}.pth\"\n        \n        if (epoch == 4 or epoch==3):\n            save_checkpoint(epoch, model, optimizer, train_trans_loss, out_path_file)\n        \n\n        test_trans_loss=test_step(model=model, \n                                dataloader=test_dataloader,\n                                loss_fn=loss_fn)\n        \n        \n        \n        \n        \n        print(f\" Epoch: {epoch},Train Translation Loss : {train_trans_loss}\")\n        print(f\" Epoch: {epoch},Test Translation Loss : {test_trans_loss}\")\n        #scheduler.step(test_trans_loss)\n        \n\n        results[\"Train Translation Loss\"].append(train_trans_loss)\n        results[\"Test Translation Loss\"].append(test_trans_loss)\n\n        metrics = {\"epoch\": epoch, \n                   \"Train Translation Loss\": train_trans_loss,\n                   \"Test Translation Loss\" :  test_trans_loss\n                  }\n        \n        wandb.log(metrics)\n        \n    return results\n                         ","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.148399Z","iopub.status.idle":"2023-06-12T15:56:42.148941Z","shell.execute_reply.started":"2023-06-12T15:56:42.148764Z","shell.execute_reply":"2023-06-12T15:56:42.148783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### Model Tester\n#train_step(model,train_dataloader,loss_fn,optimizer)\n#test_step(model,test_dataloader,loss_fn)","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.149995Z","iopub.status.idle":"2023-06-12T15:56:42.150784Z","shell.execute_reply.started":"2023-06-12T15:56:42.150589Z","shell.execute_reply":"2023-06-12T15:56:42.150608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"'''for epoch in tqdm(range(0,100), total=EPOCHS, desc='Epochs', position=0, leave=True):\n    train_rot_loss= train_step(model=model,\n                               dataloader=train_dataloader,\n                               loss_fn=loss_fn,\n                               optimizer=optimizer)\n    test_rot_loss=test_step(model=model, \n                            dataloader=test_dataloader,\n                            loss_fn=loss_fn)\n    print(f\" Epoch: {epoch} Train Rotation Loss : {train_rot_loss}\")\n    print(f\" Epoch: {epoch} Test Rotation Loss : {test_rot_loss}\")'''","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.151868Z","iopub.status.idle":"2023-06-12T15:56:42.152216Z","shell.execute_reply.started":"2023-06-12T15:56:42.152053Z","shell.execute_reply":"2023-06-12T15:56:42.152069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# WANDB","metadata":{}},{"cell_type":"code","source":"wandb.init(\n        project=WANDB_PROJECT,\n        config={\n            \"epochs\": 10,\n            \"batch_size\": 256,\n            \"lr\": 1e-3,\n            \"dropout\": 0.3,\n            })\n    \nsuper_loop=500\n\nfor i in range(super_loop):\n\n    df_train_pairs=df_train_pairs_original.sample(n=100)\n    df_test_pairs=df_test_pairs_original.sample(n=20)\n    train_dataset=PairDataset(df_train_pairs, image_transform,feature_transform)\n    test_dataset=PairDataset(df_test_pairs, image_transform,feature_transform)\n    train_dataloader=DataLoader(dataset=train_dataset, batch_size=BATCH_SIZE, shuffle= True)\n    test_dataloader=DataLoader(dataset=test_dataset, batch_size=BATCH_SIZE, shuffle= False)\n    \n    results=train(model=model,\n                  train_dataloader=train_dataloader, \n                  test_dataloader=test_dataloader,\n                  loss_fn=loss_fn,\n                  optimizer=optimizer,\n                  epochs=EPOCHS,\n                  out_path=out_path\n                 )","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.153269Z","iopub.status.idle":"2023-06-12T15:56:42.153602Z","shell.execute_reply.started":"2023-06-12T15:56:42.153439Z","shell.execute_reply":"2023-06-12T15:56:42.153454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2023-06-12T15:56:42.154828Z","iopub.status.idle":"2023-06-12T15:56:42.155172Z","shell.execute_reply.started":"2023-06-12T15:56:42.155008Z","shell.execute_reply":"2023-06-12T15:56:42.155024Z"},"trusted":true},"execution_count":null,"outputs":[]}]}