{"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":"gpu","dataSources":[{"sourceId":6799,"databundleVersionId":4225553,"sourceType":"competition"}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torchvision\nimport cv2\nimport os\nimport random as rand\n# import torchlens as tl\nfrom torch import nn, utils, optim\nfrom torch.nn import functional as F\nfrom torch.utils import data\nfrom torchvision import models as vis_models, datasets\nfrom torchvision.transforms import v2\nfrom glob import glob\nfrom tqdm import tqdm\nfrom sklearn.model_selection import train_test_split\ntorch.cuda.is_available()","metadata":{"execution":{"iopub.status.busy":"2024-06-20T18:15:31.123417Z","iopub.execute_input":"2024-06-20T18:15:31.123752Z","iopub.status.idle":"2024-06-20T18:15:37.489425Z","shell.execute_reply.started":"2024-06-20T18:15:31.123723Z","shell.execute_reply":"2024-06-20T18:15:37.48844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_dir = glob('/kaggle/input/imagenet-object-localization-challenge/ILSVRC/Data/CLS-LOC/train/**/*.JPEG')\ndef is_valid(dir):\n    img = cv2.imread(dir)\n    try:\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        return True\n    except:\n        return False\n\ntrain_dir, test_dir = train_test_split(\n    img_dir, test_size=0.2\n)\n\nlen(train_dir), len(test_dir)","metadata":{"execution":{"iopub.status.busy":"2024-06-20T18:15:37.491378Z","iopub.execute_input":"2024-06-20T18:15:37.492105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ImageDataset(data.Dataset):\n    def __init__(self, image_dir):\n        self.img_dict = {}\n        for img_dir in image_dir:\n            img_module = img_dir.split('/')[-2]\n            if img_module in list(self.img_dict.keys()):\n                self.img_dict[img_module].append(img_dir)\n            else:\n                self.img_dict[img_module]=[img_dir]\n        self.img_modules = list(self.img_dict.keys())\n        self.T = v2.Compose([\n            v2.ToTensor(),\n            v2.ToDtype(torch.float32),\n            v2.Resize(256),\n            v2.CenterCrop(224),\n            v2.Normalize(mean = [0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ])\n        self.len_dir = len(image_dir)\n        self.prev_img = {}\n        \n    def __len__(self, ):\n        return self.len_dir\n    \n    def read_img(self, img_dir):\n        img = cv2.imread(img_dir)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        return self.T(img)\n    \n    def __getitem__(self, idx):\n        pos_module = rand.choice(self.img_modules)\n        pos_path = rand.choice(self.img_dict[pos_module])\n        pos_img = self.read_img(pos_path)\n        while True:\n            neg_module = rand.choice(self.img_modules)\n            if neg_module!=pos_module:\n                neg_path = rand.choice(self.img_dict[neg_module])\n                neg_img = self.read_img(neg_path)\n                while True:\n                    anc_path = rand.choice(self.img_dict[pos_module])\n                    if anc_path!=pos_path:\n                        break\n                anc_img = self.read_img(anc_path)\n                break\n        return pos_img, anc_img, neg_img","metadata":{"execution":{"iopub.status.busy":"2024-06-20T08:41:49.420633Z","iopub.execute_input":"2024-06-20T08:41:49.421297Z","iopub.status.idle":"2024-06-20T08:41:49.498789Z","shell.execute_reply.started":"2024-06-20T08:41:49.42127Z","shell.execute_reply":"2024-06-20T08:41:49.497812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResLink(nn.Module):\n    def __init__(self, in_ch) -> None:\n        super(ResLink, self).__init__()\n        self.con1 = nn.Conv2d(in_ch, in_ch*2, kernel_size=(3,3), stride=(1,1), padding=(1,1), bias=True)\n        self.btn = nn.BatchNorm2d(2*in_ch)\n    \n    def forward(self, x):\n        x = self.con1(x)\n        x = self.btn(x)\n        return x\n        \nclass CNNBlock(nn.Module):\n    def __init__(self, in_ch) -> None:\n        super(CNNBlock, self).__init__()\n        self.con1_1 = nn.Conv2d(in_ch, in_ch*2, kernel_size=(3,3), stride=(1,1), padding=(1,1), bias=False)\n        self.btn1_1 = nn.BatchNorm2d(in_ch*2)\n        self.rel1_1 = nn.ReLU()\n        self.con2_1 = nn.Conv2d(in_ch*2, in_ch*2, kernel_size=(3,3), stride=(1,1), padding=(1,1), bias=False)\n        self.btn2_1 = nn.BatchNorm2d(in_ch*2)\n        self.rel2_1 = nn.ReLU()\n        self.res_link = ResLink(in_ch)\n        self.downsample = nn.Sequential(\n            nn.Conv2d(in_ch*4, out_channels=in_ch*4, kernel_size=(3,3), stride=(2,2)),\n            nn.BatchNorm2d(in_ch*4),\n        )\n\n    def forward(self, x):\n        x1 = self.con1_1(x)\n        x1 = self.btn1_1(x1)\n        x1 = self.rel1_1(x1)\n        x1 = self.con2_1(x1)\n        x1 = self.btn2_1(x1)\n        x1 = self.rel2_1(x1)\n        x2 = self.res_link(x)\n        x = torch.cat([x1, x2], dim=1)\n        x = self.downsample(x)\n        return x\n\nclass EncoderCNN(nn.Module):\n    def __init__(self, in_ch, out_in_ch) -> None:\n        super(EncoderCNN, self).__init__()\n        self.con_in = nn.Conv2d(in_ch, out_channels=out_in_ch, kernel_size=(3,3), stride=(1,1))\n        self.btn_in = nn.BatchNorm2d(32)\n        self.rel_in = nn.ReLU()\n        self.cnn1 = CNNBlock(out_in_ch)\n        self.cnn2 = CNNBlock(out_in_ch*4)\n        # self.cnn3 = CNNBlock(out_in_ch*16)\n        self.avg = nn.AdaptiveAvgPool2d((7,7))\n        self.out = nn.Sequential(\n            nn.Linear(in_features=25088, out_features=1024),\n            nn.Dropout1d(),\n            nn.Linear(in_features = 1024, out_features = 128)\n        )\n    \n    def forward(self, img):\n        x = self.con_in(img)\n        x = self.btn_in(x)\n        x = self.rel_in(x)\n        x = self.cnn1(x)\n        x = self.cnn2(x)\n        # x = self.cnn3(x)\n        x = self.avg(x)\n        bs, ch, hi, ws = x.shape\n        x = x.reshape(bs, -1)\n        # print(x.shape)\n        x = self.out(x)\n        return x\nimg_enc = EncoderCNN(3, 32).to('cuda')","metadata":{"execution":{"iopub.status.busy":"2024-06-20T08:43:39.992325Z","iopub.execute_input":"2024-06-20T08:43:39.993062Z","iopub.status.idle":"2024-06-20T08:43:40.323276Z","shell.execute_reply.started":"2024-06-20T08:43:39.993032Z","shell.execute_reply":"2024-06-20T08:43:40.3223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"adam = optim.Adam(img_enc.parameters(), lr=3e-5)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(adam, 'min')\ncriterion = nn.TripletMarginWithDistanceLoss()\nn_epochs = 50\ntrain_dataset = data.DataLoader(ImageDataset(train_dir), batch_size=4, shuffle=True)\nval_dataset = data.DataLoader(ImageDataset(test_dir), batch_size=4)","metadata":{"execution":{"iopub.status.busy":"2024-06-20T08:43:44.485718Z","iopub.execute_input":"2024-06-20T08:43:44.486071Z","iopub.status.idle":"2024-06-20T08:44:22.922197Z","shell.execute_reply.started":"2024-06-20T08:43:44.486044Z","shell.execute_reply":"2024-06-20T08:44:22.92142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dist_func = nn.PairwiseDistance()\nfor epoch in range(n_epochs):\n    step = 0\n    tot_loss = 0\n    same_dist = 0\n    nsame_dist = 0\n    img_enc.train()\n    train_progress = tqdm(train_dataset, desc=f\"EPOCHS:{epoch+1} ||\")\n    # test_progress = tqdm(val_dataset, desc=\"Val: \")\n    for pos, anc, neg in train_progress:\n        adam.zero_grad()\n        step+=1\n        out_pos, out_anc, out_neg = img_enc(pos.to('cuda')), img_enc(anc.to('cuda')), img_enc(neg.to('cuda'))\n        loss = criterion(out_anc, out_pos, out_neg)\n        same = dist_func(out_anc, out_pos)\n        nsame = dist_func(out_anc, out_neg)\n        train_progress.set_postfix({\n            'Loss':loss.cpu().item(),\n            'Same Dist':same.cpu().mean().item(),\n            'Not Same Dist':nsame.cpu().mean().item()\n        })\n        tot_loss+=loss.cpu().item()\n        same_dist+=same.cpu().mean().item()\n        nsame_dist+=nsame.cpu().mean().item()\n        loss.backward()\n        adam.step()\n    print(f'|| Avg Loss:{tot_loss/step} || Avg Same Dist:{same_dist/step}||Avg Not Same Dist:{nsame_dist/step}')","metadata":{"execution":{"iopub.status.busy":"2024-06-20T08:44:22.923593Z","iopub.execute_input":"2024-06-20T08:44:22.923875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}