{"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":"none","dataSources":[{"sourceId":6818,"databundleVersionId":1960702,"sourceType":"competition"}],"dockerImageVersionId":30646,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport os\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset,DataLoader,random_split\nfrom torchvision.utils import make_grid\nfrom torchvision import transforms\nimport torchvision\nfrom tqdm import tqdm\nfrom skimage.transform import resize\nimport numpy as np\nimport cv2\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-02-01T15:02:53.389613Z","iopub.execute_input":"2024-02-01T15:02:53.390118Z","iopub.status.idle":"2024-02-01T15:02:53.398185Z","shell.execute_reply.started":"2024-02-01T15:02:53.390082Z","shell.execute_reply":"2024-02-01T15:02:53.396908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_whale_by_id(root_path,\n                          df : pd.DataFrame,\n                          whale_id,\n                          num_samples):\n    class_name = df.Id.unique()[whale_id]\n    rows = df[df.Id == class_name].iloc[:num_samples]\n    imgs = []\n    for row in rows.iterrows():\n        img = cv2.imread(os.path.join(root_path,row[1]['Image']))\n        img = cv2.cvtColor(img,cv2.COLOR_BGR2RGB)\n        try:\n            if img.ndim == 3:\n                img = resize(img,(224,224)).transpose(2,0,1)\n            elif img.ndim == 2:\n                print('Exception')\n                img = resize(img,(224,224))[None,:,:]\n                img = np.repeat(3,1,1)\n        except:\n            print(img.shape)\n            input()\n        imgs.append(img)\n    imgs = torch.from_numpy(np.stack(imgs))\n    imgs_grid = make_grid(imgs,nrow = 4)\n    plt.title(class_name)\n    plt.imshow(imgs_grid.permute(1,2,0))\n    plt.axis(\"off\")\n    plt.show()\n        ","metadata":{"execution":{"iopub.status.busy":"2024-02-01T14:25:51.017557Z","iopub.execute_input":"2024-02-01T14:25:51.018502Z","iopub.status.idle":"2024-02-01T14:25:51.030853Z","shell.execute_reply.started":"2024-02-01T14:25:51.018461Z","shell.execute_reply":"2024-02-01T14:25:51.029388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"root_path = '/kaggle/input/humpback-whale-identification/train'\ncsv_path = '/kaggle/input/humpback-whale-identification/train.csv'\ndf = pd.read_csv(csv_path)\n\nfor i in range(len(df.Id.unique())):\n    visualize_whale_by_id(root_path,df,i,16)","metadata":{"execution":{"iopub.status.busy":"2024-02-01T14:26:19.747343Z","iopub.execute_input":"2024-02-01T14:26:19.747745Z","iopub.status.idle":"2024-02-01T14:26:46.133509Z","shell.execute_reply.started":"2024-02-01T14:26:19.747713Z","shell.execute_reply":"2024-02-01T14:26:46.131653Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i,name in enumerate(df.Id.unique()):\n    if name == 'new_whale':\n        visualize_whale_by_id(root_path,df,i,32)","metadata":{"execution":{"iopub.status.busy":"2024-02-01T09:03:31.507904Z","iopub.execute_input":"2024-02-01T09:03:31.509185Z","iopub.status.idle":"2024-02-01T09:03:31.513779Z","shell.execute_reply.started":"2024-02-01T09:03:31.509146Z","shell.execute_reply":"2024-02-01T09:03:31.512703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    root_path = '/kaggle/input/humpback-whale-identification/train'\n    csv_path = '/kaggle/input/humpback-whale-identification/train.csv'\n    total_epoch = 100\n    lr = 0.001\n    save_dir = '/kaggle/working/weighted'\n    alpha = 0.1\n    batch_size = 16\n    device = 'cpu'","metadata":{"execution":{"iopub.status.busy":"2024-02-01T15:06:00.808598Z","iopub.execute_input":"2024-02-01T15:06:00.809028Z","iopub.status.idle":"2024-02-01T15:06:00.815213Z","shell.execute_reply.started":"2024-02-01T15:06:00.808996Z","shell.execute_reply":"2024-02-01T15:06:00.813884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TripletDataset(Dataset):\n    def __init__(self,root_path,df,transform_aug,transform_normal):\n        super().__init__()\n        self.trip_list = []\n        self.transform_aug  = transform_aug\n        self.transform_normal = transform_normal\n        \n        for i in tqdm(range(len(df))):\n            img_name,label = df.iloc[i]['Image'],df.iloc[i]['Id']\n            anchor = os.path.join(root_path,img_name)\n            pos_idxs = df[(df.Id == label) & (df.Image != img_name)]\n            neg_idxs = df[df.Id != label]\n            \n            neg_idx = np.random.randint(0,len(neg_idxs))\n            \n            if len(pos_idxs) == 0:\n                pos = -1\n                neg = os.path.join(root_path,df.iloc[neg_idx]['Image'])\n            else:\n                pos_idx = np.random.randint(0,len(pos_idxs))\n                neg = os.path.join(root_path,df.iloc[neg_idx]['Image'])\n                pos = os.path.join(root_path,df.iloc[pos_idx]['Image'])\n                \n            self.trip_list.append((anchor,pos,neg))\n            \n    def __len__(self):\n        return len(self.trip_list)\n    def __getitem__(self,idx):\n        anchor_path,pos_path,neg_path = self.trip_list[idx]\n        anchor = cv2.imread(anchor_path)\n        neg = cv2.imread(neg_path)\n        if pos_path == -1:\n            pos = self.transform_aug(anchor)\n        else:\n            pos = cv2.imread(pos_path)\n            pos = self.transform_normal(pos)\n        anchor = self.transform_normal(anchor)\n        neg = self.transform_normal(neg)\n        return anchor,pos,neg\n            \n        \n\ntransform_aug = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.RandomResizedCrop((224,224),antialias=True),\n    transforms.RandomVerticalFlip(p = 1.0),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\ntransform_normal = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Resize((224,224)),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\nargs = CFG()\ndf = pd.read_csv(args.csv_path)\ndataset = TripletDataset(args.root_path,df,transform_aug,transform_normal)\nlen(dataset)","metadata":{"execution":{"iopub.status.busy":"2024-02-01T15:06:00.951384Z","iopub.execute_input":"2024-02-01T15:06:00.951816Z","iopub.status.idle":"2024-02-01T15:06:09.456209Z","shell.execute_reply.started":"2024-02-01T15:06:00.951783Z","shell.execute_reply":"2024-02-01T15:06:09.454223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(dataset)","metadata":{"execution":{"iopub.status.busy":"2024-02-01T15:06:14.10613Z","iopub.execute_input":"2024-02-01T15:06:14.106553Z","iopub.status.idle":"2024-02-01T15:06:14.114857Z","shell.execute_reply.started":"2024-02-01T15:06:14.106515Z","shell.execute_reply":"2024-02-01T15:06:14.113117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TripletLoss(nn.Module):\n    def __init__(self,alpha):\n        super().__init__()\n        self.alpha = alpha\n    def forward(self,anchor,pos,neg):\n        \"\"\"\n        anchor : [N,1024]\n        pos : [N,1024]\n        neg : [N,1024]\n        \"\"\"\n        p = torch.linalg.norm(anchor-pos)\n        n = torch.linalg.norm(anchor - neg)\n        return torch.clamp(p - (self.alpha - n),min = 0.0).mean()\n        ","metadata":{"execution":{"iopub.status.busy":"2024-02-01T15:06:09.457334Z","iopub.status.idle":"2024-02-01T15:06:09.457756Z","shell.execute_reply.started":"2024-02-01T15:06:09.457554Z","shell.execute_reply":"2024-02-01T15:06:09.45757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyNetwork(nn.Module):\n    def __init__(self,out_features):\n        super().__init__()\n        self.backbone = torchvision.models.resnet50(pretrained = True)\n        self.backbone.fc = nn.Sequential(\n            nn.Linear(2048,1024),\n            nn.BatchNorm1d(1024),\n            nn.ReLU(),\n            nn.Linear(1024,out_features)\n        )\n    def forward(self,imgs):\n        return self.backbone(imgs)\n        ","metadata":{"execution":{"iopub.status.busy":"2024-02-01T15:06:09.459239Z","iopub.status.idle":"2024-02-01T15:06:09.459772Z","shell.execute_reply.started":"2024-02-01T15:06:09.459537Z","shell.execute_reply":"2024-02-01T15:06:09.459564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_epoch(cfg,epoch,model,train_loader,optimizer,criterion):\n    model.train()\n    losses = 0.0\n    for idx,(anchor,pos,neg) in enumerate(train_loader):\n        anchor,pos,neg = anchor.to(cfg.device),pos.to(cfg.device),neg.to(cfg.device)\n        imgs = torch.concat([anchor,pos,neg],dim = 0)\n        \n        imgs_feat = model(imgs)\n        a_feat,pos_feat,neg_feat = torch.split(imgs_feat,imgs_feat.shape[0] // 3,dim = 0)\n        \n        optimizer.zero_grad()\n        loss = criterion(a_feat,pos_feat,neg_feat)\n        loss.backward()\n        optimizer.step()\n        \n        lr = optimizer.param_groups[0]['lr']\n        if idx % 5 == 0:\n            print(f\"Epoch [{epoch}/{cfg.total_epoch}] | Step [{idx}/{len(train_loader)}] | Lr : {lr} | Loss : {loss:.4f}\")\n        losses += loss.detach().cpu().item()\n    print(f\"Epoch [{epoch}/{cfg.total_epoch}] | Avg_loss : {losses/len(train_loader):.4f}\")\n    return losses / len(train_loader)\n            \ndef val_epoch(cfg,epoch,model,test_loader,criterion):\n    model.eval()\n    losses = 0.0\n    with torch.no_grad():\n        for idx,(anchor,pos,neg) in enumerate(test_loader):\n            anchor,pos,neg = anchor.to(cfg.device),pos.to(cfg.device),neg.to(cfg.device)\n            imgs = torch.concat([anchor,pos,neg],dim = 0)\n\n            imgs_feat = model(imgs)\n            a_feat,pos_feat,neg_feat = torch.split(imgs_feat,3,dim = 0)\n\n            loss = criterion(a_feat,pos_feat,neg_feat)\n            losses += loss.detach().cpu().item()\n        print(f\"Test Epoch [{epoch}/{cfg.total_epoch}] | Avg_loss : {losses/len(test_loader):.4f}\")\n    return losses / len(train_loader)\n    ","metadata":{"execution":{"iopub.status.busy":"2024-02-01T15:10:16.296137Z","iopub.execute_input":"2024-02-01T15:10:16.29662Z","iopub.status.idle":"2024-02-01T15:10:16.312097Z","shell.execute_reply.started":"2024-02-01T15:10:16.296586Z","shell.execute_reply":"2024-02-01T15:10:16.310542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = CFG()\ntrain_losses = []\ntest_losses = []\nbest_val_loss = -1\nmodel = MyNetwork(512)\ncriterion =TripletLoss(alpha = cfg.alpha) \noptimizer = torch.optim.Adam(model.parameters(),lr = cfg.lr)\n\n\ntrain_size = int(0.8 * len(dataset))\ntest_size = len(dataset) - train_size\n\ntrain_dataset, test_dataset = random_split(dataset, [train_size, test_size])\n\n# Now you can create DataLoader for training and testing\ntrain_loader = DataLoader(train_dataset, batch_size=cfg.batch_size, shuffle=True)\ntest_loader = DataLoader(test_dataset, batch_size=cfg.batch_size, shuffle=False)\n\nfor epoch in range(cfg.total_epoch):\n    train_loss = train_epoch(cfg,epoch,model,train_loader,optimizer,criterion)\n    test_loss = val_epoch(cfg,epoch,model,test_loader)\n    train_losses.append(train_loss)\n    test_losses.append(test_loss)\n    if test_loss > best_val_loss:\n        best_val_loss = test_loss\n        torch.save(model.state_dict(),f'model_{epoch}.pt')","metadata":{"execution":{"iopub.status.busy":"2024-02-01T15:10:16.457065Z","iopub.execute_input":"2024-02-01T15:10:16.457495Z","iopub.status.idle":"2024-02-01T15:12:26.189299Z","shell.execute_reply.started":"2024-02-01T15:10:16.457463Z","shell.execute_reply":"2024-02-01T15:12:26.187137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"a = torch.randn(10,50)\nlen(torch.split(a,5,dim = 0))","metadata":{"execution":{"iopub.status.busy":"2024-02-01T15:12:26.191736Z","iopub.status.idle":"2024-02-01T15:12:26.192217Z","shell.execute_reply.started":"2024-02-01T15:12:26.192Z","shell.execute_reply":"2024-02-01T15:12:26.192019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}