{"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 torch.nn.functional as F\nfrom torch.nn import Parameter\nfrom torch import nn\nimport torch\nimport math\nfrom torchvision.models import resnet50\nclass skip1(nn.Module):\n    def __init__(self, in_channels, out_channels, stride, expansion):\n        super(skip1, self).__init__()\n        \n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0, bias=False)\n        torch.nn.init.xavier_uniform(self.conv1.weight.data)\n        self.batch_norm1 = nn.BatchNorm2d(out_channels)\n        \n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)\n        torch.nn.init.xavier_uniform(self.conv2.weight.data)\n        self.batch_norm2 = nn.BatchNorm2d(out_channels)\n        \n        self.conv3 = nn.Conv2d(out_channels, out_channels*expansion, kernel_size=1, stride=1, padding=0, bias=False)\n        torch.nn.init.xavier_uniform(self.conv3.weight.data)\n        self.batch_norm3 = nn.BatchNorm2d(out_channels*expansion)\n        \n        self.conv4 = nn.Conv2d(in_channels, out_channels*expansion, kernel_size=1, stride=stride, padding=0, bias=False)\n        torch.nn.init.xavier_uniform(self.conv4.weight.data)\n        self.batch_norm4 = nn.BatchNorm2d(out_channels*expansion)\n        \n        self.relu = nn.ReLU(inplace=True)\n        \n    def forward(self, x):\n        identity = x.clone()\n        x = self.batch_norm1(self.conv1(x))\n        \n        x = self.batch_norm2(self.conv2(x))\n        \n        x = self.batch_norm3(self.conv3(x))\n        \n        identity = self.batch_norm4(self.conv4(identity))\n        \n        x+=identity\n        \n        return x\n\nclass skip2(nn.Module):\n    def __init__(self, in_channels, out_channels, expansion):\n        super(skip2, self).__init__()\n        \n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0, bias=False)\n        torch.nn.init.xavier_uniform(self.conv1.weight.data)\n        self.batch_norm1 = nn.BatchNorm2d(out_channels)\n        \n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)\n        torch.nn.init.xavier_uniform(self.conv2.weight.data)\n        self.batch_norm2 = nn.BatchNorm2d(out_channels)\n        \n        self.conv3 = nn.Conv2d(out_channels, out_channels*expansion, kernel_size=1, stride=1, padding=0, bias=False)\n        torch.nn.init.xavier_uniform(self.conv3.weight.data)\n        self.batch_norm3 = nn.BatchNorm2d(out_channels*expansion)\n        \n        self.relu = nn.ReLU(inplace=True)\n        \n    def forward(self, x):\n        identity = x.clone()\n        x = self.batch_norm1(self.conv1(x))\n        \n        x = self.batch_norm2(self.conv2(x))\n        \n        x = self.batch_norm3(self.conv3(x))\n        \n        x+=identity\n        \n        return self.relu(x)\n        \n\nclass ResNet50(nn.Module):\n    def __init__(self, num_classes, num_channels=3):\n      super(ResNet50, self).__init__()\n      self.conv1 = nn.Conv2d(num_channels, 64, kernel_size=7, stride=2, padding=3, bias=False)\n      torch.nn.init.xavier_uniform(self.conv1.weight.data)\n      self.batch_norm1 = nn.BatchNorm2d(64)\n      self.relu = nn.ReLU(inplace=True)\n      self.max_pool = nn.MaxPool2d(kernel_size = 3, stride=2, padding=1)\n\n      self.layer1 =  nn.Sequential(skip1(64, 64, 1, 4)\n      ,skip2(256, 64, 4)\n      ,skip2(256, 64, 4))\n      \n      self.layer2 =  nn.Sequential(skip1(256, 128, 2, 4)\n      ,skip2(512, 128, 4)\n      ,skip2(512, 128, 4)\n      ,skip2(512, 128, 4))\n      \n      self.layer3 =  nn.Sequential(skip1(512, 256, 2, 4)\n      ,skip2(1024, 256, 4)\n      ,skip2(1024, 256, 4)\n      ,skip2(1024, 256, 4)\n      ,skip2(1024, 256, 4)\n      ,skip2(1024, 256, 4))\n\n      \n      self.layer4 =  nn.Sequential(skip1(1024, 512, 2, 4)\n      ,skip2(2048, 512, 4)\n      ,skip2(2048, 512, 4))\n    def forward(self, x):\n        a = self.batch_norm1(self.conv1(x))\n        b = self.max_pool(a)\n\n        c = self.layer1(b)\n        d = self.layer2(c)\n        l3 = self.layer3(d)\n        l4 = self.layer4(l3)\n        \n        return l4,l3\ndef get_pretrained_resnet50(mod):\n  # mod=ResNet50(100)\n  m1=mod.state_dict()\n  m2=resnet50(pretrained=True).state_dict()\n  for a,b in zip(list(m1),list(m2)):\n    m1[a]=m2[b]\n  mod.load_state_dict(m1)\n  return mod\n\n\n\nclass local(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(local, self).__init__()\n        self.in_channels=in_channels\n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=6, dilation=6, bias=False)\n        torch.nn.init.xavier_uniform(self.conv1.weight.data)\n        self.conv2 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=12, dilation=12, bias=False)\n        torch.nn.init.xavier_uniform(self.conv2.weight.data)\n        self.conv3 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=18, dilation=18, bias=False)\n        torch.nn.init.xavier_uniform(self.conv3.weight.data)\n        self.conv4 = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0, bias=False)\n        torch.nn.init.xavier_uniform(self.conv4.weight.data)\n        self.relu = nn.ReLU(inplace=True)\n\n        self.conv5 = nn.Conv2d(out_channels*4, in_channels, kernel_size=1, stride=1, padding=0, bias=False)\n        torch.nn.init.xavier_uniform(self.conv5.weight)\n        \n    def forward(self, x):\n        shape=x.shape\n        a = self.conv1(x)\n        b = self.conv2(x)\n        c = self.conv3(x)\n        \n        gap=nn.functional.interpolate(self.relu(self.conv4(torch.mean(x,[2,3]).view(-1,self.in_channels,1,1))),scale_factor=[shape[2],shape[3]], mode=\"bilinear\")\n        conc=torch.concat([a, b, c, gap],1)\n        return self.conv5(conc)\n        \nclass SA(nn.Module):\n    def __init__(self, in_channels):\n        super(SA, self).__init__()\n        \n        self.conv1 = nn.Conv2d(in_channels, 1024, kernel_size=1, bias=False)\n        torch.nn.init.xavier_uniform(self.conv1.weight.data)\n        self.bn1 = nn.BatchNorm2d(1024)\n        self.relu = nn.ReLU(inplace=True)\n\n        self.conv2 = nn.Conv2d(1024, 1024, kernel_size=1, bias=False)\n        torch.nn.init.xavier_uniform(self.conv2.weight.data)\n        self.sp = torch.nn.Softplus()\n        \n    def forward(self, x):\n        x = self.bn1(self.conv1(x))\n        a = x.clone()\n        b = self.sp(self.conv2(a))\n        a = torch.norm(a,1,1)\n        return torch.unsqueeze(a,1)*b\n\nclass OFM(nn.Module):\n    def __init__(self):\n        super(OFM, self).__init__()\n        self.convert = nn.Linear(2048, 1024)\n        torch.nn.init.xavier_uniform(self.convert.weight.data)\n        self.convert.bias.data.zero_()\n        \n    def forward(self, l, g):\n        g = self.convert(g)\n        shape=l.shape\n        \n        g_norm=torch.norm(g,1,1).reshape(-1,1,1,1)\n        mid=torch.bmm(g.unsqueeze(1),torch.flatten(l.view(shape[0],shape[1],-1),start_dim=2))\n        out=torch.bmm(g.unsqueeze(2),mid).view(l.shape)\n        out=out/(g_norm*g_norm)\n        return l-out\n\nclass Arcface(nn.Module):\n    def __init__(self, out_features, in_features, device, scale=30, margin=0.15):\n        super(Arcface, self).__init__()\n        self.in_features = in_features\n        self.device = device\n        self.out_features = out_features\n        self.scale = scale\n        self.margin = margin\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.cos_m = math.cos(margin)\n        self.sin_m = math.sin(margin)\n        self.th = math.cos(math.pi - margin)\n        self.mm = math.sin(math.pi - margin) * margin\n\n    def forward(self, input, label):\n        # --------------------------- cos(theta) & phi(theta) ---------------------------\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n        phi = cosine * self.cos_m - sine * self.sin_m\n        \n        phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        \n        one_hot = torch.zeros(cosine.size(), device=self.device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        \n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.scale\n\n        return output\n\n\nclass GeM(nn.Module):\n    def __init__(self, p=3):\n        super(GeM, self).__init__()\n        self.p=p\n        \n    def forward(self, x):\n        a=torch.pow(torch.clip(x,1e-6),self.p)\n        b=torch.mean(a,dim=[2,3])\n        c=torch.pow(b,1/self.p)\n        return c, b, a\n    \nclass DOLG(nn.Module):\n  def __init__(self,n_class, device,L3=1024,L4=2048):\n    super(DOLG, self).__init__()\n    self.r50=ResNet50(100)\n    self.r50=get_pretrained_resnet50(self.r50)\n    self.gem=GeM()\n    self.l=local(1024,512)\n    self.a=SA(1024)\n    self.orthogonal=OFM()\n    self.pred=nn.Linear(L3+L4, 512)\n    torch.nn.init.xavier_uniform(self.pred.weight.data)\n    self.pred.bias.data.zero_()\n    self.arc=Arcface(n_class,512,device)\n  def forward(self,x,y,mode='test'):\n    G,L=self.r50(x)\n    \n    G,_,_=self.gem(G)\n    L=self.l(L)\n    L=self.a(L)\n    L=torch.mean(self.orthogonal(L,G),dim=[2,3])\n    all_features=torch.concat([G,L],-1)\n    sout=self.pred(all_features)\n    return self.arc(sout,y)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-19T03:39:15.763586Z","iopub.execute_input":"2022-12-19T03:39:15.763926Z","iopub.status.idle":"2022-12-19T03:39:18.125241Z","shell.execute_reply.started":"2022-12-19T03:39:15.763850Z","shell.execute_reply":"2022-12-19T03:39:18.124003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install wandb","metadata":{"execution":{"iopub.status.busy":"2022-12-19T03:39:18.127333Z","iopub.execute_input":"2022-12-19T03:39:18.127868Z","iopub.status.idle":"2022-12-19T03:39:29.319760Z","shell.execute_reply.started":"2022-12-19T03:39:18.127832Z","shell.execute_reply":"2022-12-19T03:39:29.318393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.utils.data as data\nimport cv2\nclass GLDV2(data.Dataset):\n    def __init__(self, df, dtype, csize=256, rsize=300) -> None:\n        self.df=df\n        self.dtype=dtype\n        self.rsize=rsize\n        self.csize=csize\n    def __len__(self) -> int:\n        return self.df.shape[0]\n\n    def __getitem__(self, index: int):\n        row = self.df.loc[index]\n        img_id = row['id']\n        label = row['landmark_id']\n        zero = row['0']\n        one = row['1']\n        two = row['2']\n        path='/kaggle/input/landmark-retrieval-2021/train/'+str(zero)+'/'+str(one)+'/'+str(two)+'/'+img_id+'.jpg'\n        img=cv2.resize(cv2.imread(path),(self.rsize,self.rsize), interpolation = cv2.INTER_AREA).astype('float16')\n        img = img[:, :, (2, 1, 0)]\n        if self.dtype=='train':\n          t1 = RandomCrop((self.csize, self.csize))\n          t2 = Normalize((0.485*255, 0.456*255, 0.406*255), (0.229*255, 0.224*255, 0.225*255))\n          return t2(t1(torch.from_numpy(img).permute(2, 0, 1))), torch.from_numpy(np.asarray([label]))\n        else:\n          return torch.from_numpy(img).permute(2, 0, 1), torch.from_numpy(np.asarray([label]))","metadata":{"execution":{"iopub.status.busy":"2022-12-19T03:39:29.322218Z","iopub.execute_input":"2022-12-19T03:39:29.323301Z","iopub.status.idle":"2022-12-19T03:39:29.498862Z","shell.execute_reply.started":"2022-12-19T03:39:29.323256Z","shell.execute_reply":"2022-12-19T03:39:29.497956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom tqdm import tqdm\nfrom torch import nn, optim\nimport numpy as np\nfrom sklearn.metrics import accuracy_score\nfrom statistics import *\nimport wandb\nimport os\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nsecret_value_0 = user_secrets.get_secret(\"your-secret-label\")\nos.environ['WANDB_API_KEY'] = secret_value_0\nwandb.login()\nclass Trainingrun:\n    \"\"\"\n    Custom model training\n    \"\"\"\n\n    def __init__(self, train_loader, valid_loader, model, last_epoch, epochs=100):\n        self.model = model\n        self.last_epoch = last_epoch\n        self.epochs=epochs\n        params = model.parameters()\n        self.device = device\n        self.optimizer = optim.SGD(params,\n                                    lr=0.027066694599366656,\n                                   momentum=0.9)\n        self.scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(self.optimizer, T_max=(100-5)*len(train_loader), eta_min=0)\n        self.criterion = nn.CrossEntropyLoss()\n        self.train_loader = train_loader\n        self.valid_loader = valid_loader\n    \n    def validate(self):\n        valid = tqdm(self.valid_loader)\n        ls=[]\n        for en,batch in enumerate(valid):\n            image, label = batch\n            image = image.to(self.device).float()\n            label = label.to(self.device)\n            logits = self.model(image, label, 'test')\n            label = torch.squeeze(label,1)\n            loss = self.criterion(logits, label).detach().cpu().numpy().item()\n            ls.append(loss)\n            del([image, label])\n            gc.collect()\n        return mean(ls)\n        \n    def train(self, log_interval: int):\n        \"\"\"\n        Training DGRec\n        \"\"\"\n        total_loss = {}\n        with wandb.init(project=\"GARec-full\", entity=\"viven-inc-object-detection\", resume=True, id='lyric-firefly-3'):\n            wandb.watch(self.model, self.criterion, log=\"all\", log_freq=10)\n            for epoch in range(self.last_epoch, self.epochs+1):\n                total_loss[epoch] = []\n                self.model.train()\n                enm = 0\n                trn = tqdm(self.train_loader)\n                for en,batch in enumerate(trn):\n\n                    if epoch<5:\n                      for param_group in self.optimizer.param_groups:\n                          param_group['lr'] = 1e-4+(((epoch-1)*len(self.train_loader))+en)*(5e-2-1e-4)/(4*len(self.train_loader))\n                    else:\n                      self.scheduler.step()\n                    image, label = batch\n                    image = image.to(self.device).float()\n                    label = label.to(self.device)\n                    logits = self.model(image, label, 'train')\n                    label = torch.squeeze(label,1)\n                    loss = self.criterion(logits, label)\n                    del([image, label])\n                    gc.collect()\n\n                    self.optimizer.zero_grad()\n                    loss.backward()\n                    self.optimizer.step()\n                    \n                    if en==0:\n                      print(\"loading weights and optimizer\")\n                      print(\"*\"*20)\n                      checkpoint = torch.load(\n                          '/kaggle/input/full-data-gldv2-256-validation-4/DOLG'+str(epoch-1)+'.pth')\n                      self.model.load_state_dict(checkpoint['model_state_dict'])\n                      self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n                      self.scheduler.load_state_dict(checkpoint['scheduler'])\n                    if en!=0:\n                        total_loss[epoch].append(loss.item())\n\n                        mean_loss = np.mean(total_loss[epoch])\n                        trn.set_description(\"Epoch: \" + str(epoch) +\n                                                \" Loss: \" + str(mean_loss)+ \"Learning rate: \"+str(self.optimizer.param_groups[0][\"lr\"]))\n                if (epoch%1)==0:\n                  print(\"Saving\")\n                  print(\"*\"*20)\n                  torch.save(\n                      {\n                          'model_state_dict': self.model.state_dict(),\n                          'optimizer_state_dict': self.optimizer.state_dict(),\n                          'scheduler': self.scheduler.state_dict(),\n                          'loss': total_loss,\n                      }, 'DOLG'+str(epoch)+'.pth')\n                  wandb.log({\n                            \"epoch\": epoch,\n                            \"batch\": en,\n                            \"loss\": mean_loss\n                        })\n                vacc=self.validate()\n                print(\"Epoch: \" + str(epoch) + \" Validation Loss: \" + str(vacc))\n            wandb.save('wandb.pth')\n        return self.optimizer, self.model, total_loss","metadata":{"execution":{"iopub.status.busy":"2022-12-19T03:39:29.501631Z","iopub.execute_input":"2022-12-19T03:39:29.502251Z","iopub.status.idle":"2022-12-19T03:39:31.471231Z","shell.execute_reply.started":"2022-12-19T03:39:29.502213Z","shell.execute_reply":"2022-12-19T03:39:31.470054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\ntrain=pd.read_csv('/kaggle/input/gldv2cleanfullsplit/train_full.csv')\nvalid=pd.read_csv('/kaggle/input/gldv2cleanfullsplit/valid_full.csv')\n\ntrain['0']=train['id'].str[0]\ntrain['1']=train['id'].str[1]\ntrain['2']=train['id'].str[2]\n\nvalid['0']=valid['id'].str[0]\nvalid['1']=valid['id'].str[1]\nvalid['2']=valid['id'].str[2]","metadata":{"execution":{"iopub.status.busy":"2022-12-19T03:39:31.474405Z","iopub.execute_input":"2022-12-19T03:39:31.474808Z","iopub.status.idle":"2022-12-19T03:39:36.036056Z","shell.execute_reply.started":"2022-12-19T03:39:31.474763Z","shell.execute_reply":"2022-12-19T03:39:36.035097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport tarfile\nfrom collections import Counter\nfrom sklearn.preprocessing import LabelEncoder\n\nfrom tqdm import tqdm\nimport torch\nfrom torch.utils.data import DataLoader\nimport pandas as pd\nimport random\ndevice = torch.device('cuda')\nmodel = DOLG(81313,device)\nmodel = model.to(device)\ntrain_set = GLDV2(train,'train')\nvalid_set = GLDV2(valid,'valid')\nimport numpy as np\nimport gc\nfrom torchvision.transforms import RandomCrop, Normalize\nfrom torch import nn, optim\n\ncur_epoch = 50\nBATCH_SIZE = 96\nNUM_WORKERS = 2\ntrain_loader = DataLoader(\n    train_set,\n    batch_size=BATCH_SIZE,\n    num_workers=NUM_WORKERS,\n    drop_last=False,\n    shuffle=True,\n)\nvalid_loader = DataLoader(\n    valid_set,\n    batch_size=32,\n    num_workers=NUM_WORKERS,\n    drop_last=False,\n    shuffle=True,\n)\n\n\nrunner = Trainingrun(train_loader, valid_loader, model, cur_epoch, epochs=cur_epoch)\nprint('started training...')\noptimizer, model, loss = runner.train(1)\nprint('completed training...')\ntorch.save(\n    {\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'loss': loss,\n    }, 'DOLG.pth')","metadata":{"execution":{"iopub.status.busy":"2022-12-19T03:39:36.037430Z","iopub.execute_input":"2022-12-19T03:39:36.038313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(\n    {\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'loss': loss,\n    }, 'DOLG.pth')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}