{"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":"markdown","source":"Configs","metadata":{}},{"cell_type":"code","source":"import pathlib\nfrom pathlib import Path\nimport torch\nimport torch.utils.data\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\nimport torch.optim as optim\nimport PIL.Image\nimport albumentations.pytorch\nimport albumentations as A\nimport cv2\nimport json\nimport matplotlib.pyplot as plt\nimport math\nfrom tqdm.notebook import tqdm\nfrom typing import List, Tuple\n!pip install timm\nimport timm\nfrom torch.utils.data import DataLoader\nimport os\nfrom PIL import Image\nimport pandas as pd\nfrom sklearn import preprocessing\nfrom torch.utils.data import Dataset","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:17:57.064373Z","iopub.execute_input":"2022-05-25T08:17:57.065041Z","iopub.status.idle":"2022-05-25T08:18:23.167958Z","shell.execute_reply.started":"2022-05-25T08:17:57.064928Z","shell.execute_reply":"2022-05-25T08:18:23.166881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LmkRetrDataset(Dataset):\n    def __init__(self, data_path, cat2name_path, image_transform):\n        \n        self.root = Path(data_path)\n        self.jsfile = cat2name_path\n        self.ext = \"jpg\"\n        self.transform = image_transform\n        \n        self._load_files()\n        self._load_classes()\n        \n    def _load_files(self):\n        self.files = sorted(list(self.root.glob(f\"*/*.{self.ext}\")))\n    \n    def _load_classes(self):\n        cat2name = self._load_json(str(self.jsfile))\n        idx2class, class2idx, data_dict = [], {}, {}\n\n        for k,v in cat2name.items(): data_dict[v] = int(k)\n        list_pair = sorted(data_dict.items(), key=lambda x: x[1])\n        for (name,cat) in list_pair:\n            class2idx[name] = cat - 1\n            idx2class.append(name)\n        \n        self.class2idx = class2idx\n        self.idx2class = idx2class\n\n    def _load_json(self, path):\n        with open(path, 'r') as jsfile:\n            data = json.load(jsfile)\n        return data\n    \n    def _load_image(self, path):\n        img = np.array(Image.open(path))\n        return img\n        \n    def __len__(self):\n        return len(self.files)\n    \n    def __getitem__(self, idx):\n        impath = self.files[idx]\n        img = self._load_image(str(impath))\n        cat = int(impath.parent.name)\n        if self.transform:\n            img = self.transform(image=img)['image']\n        return img, cat-1","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:18:27.702967Z","iopub.execute_input":"2022-05-25T08:18:27.703263Z","iopub.status.idle":"2022-05-25T08:18:27.718943Z","shell.execute_reply.started":"2022-05-25T08:18:27.703232Z","shell.execute_reply":"2022-05-25T08:18:27.717819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"GEM","metadata":{}},{"cell_type":"code","source":"class GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6, requires_grad=False):\n        super(GeM, self).__init__()\n        self.p = nn.Parameter(torch.ones(1)*p, requires_grad=requires_grad)\n        self.eps = eps\n\n    def forward(self, x):\n        return self.gem(x, p=self.p, eps=self.eps)\n\n    def gem(self, x, p=3, eps=1e-6):\n        return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(1./p)\n\n    def __repr__(self):\n        return self.__class__.__name__ + '(' + 'p=' + '{:.4f}'.format(self.p.data.tolist()[0]) + ', ' + 'eps=' + str(self.eps) + ')'","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:18:30.675908Z","iopub.execute_input":"2022-05-25T08:18:30.676191Z","iopub.status.idle":"2022-05-25T08:18:30.685481Z","shell.execute_reply.started":"2022-05-25T08:18:30.676163Z","shell.execute_reply":"2022-05-25T08:18:30.684594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"ArcFace","metadata":{}},{"cell_type":"markdown","source":"\n","metadata":{}},{"cell_type":"code","source":"class ArcFace(nn.Module):\n    def __init__(self, in_features, out_features, scale_factor=64.0, margin=0.50, criterion=None):\n        super(ArcFace, self).__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n\n        if criterion:\n            self.criterion = criterion\n        else:\n            self.criterion = nn.CrossEntropyLoss()\n\n        self.margin = margin\n        self.scale_factor = scale_factor\n\n        self.weight = nn.Parameter(\n            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        # input is not l2 normalized\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n\n        phi = cosine * self.cos_m - sine * self.sin_m\n        phi = phi.type(cosine.type())\n        phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n\n        one_hot = torch.zeros(cosine.size(), device=input.device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n\n        logit = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        logit *= self.scale_factor\n\n        loss = self.criterion(logit, label)\n\n        return loss, logit","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:18:33.456154Z","iopub.execute_input":"2022-05-25T08:18:33.456459Z","iopub.status.idle":"2022-05-25T08:18:33.470406Z","shell.execute_reply.started":"2022-05-25T08:18:33.456426Z","shell.execute_reply":"2022-05-25T08:18:33.469109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MultiAtrous(nn.Module):\n    def __init__(self, in_channel, out_channel, size, dilation_rates=[3, 6, 9]):\n        super().__init__()\n        self.dilated_convs = [\n            nn.Conv2d(in_channel, int(out_channel/4),\n                      kernel_size=3, dilation=rate, padding=rate)\n            for rate in dilation_rates\n        ]\n        self.gap_branch = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(in_channel, int(out_channel/4), kernel_size=1),\n            nn.ReLU(),\n            nn.Upsample(size=(size, size), mode='bilinear')\n        )\n        self.dilated_convs.append(self.gap_branch)\n        self.dilated_convs = nn.ModuleList(self.dilated_convs)\n\n    def forward(self, x):\n        local_feat = []\n        for dilated_conv in self.dilated_convs:\n            local_feat.append(dilated_conv(x))\n        local_feat = torch.cat(local_feat, dim=1)\n        return local_feat","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:18:37.276126Z","iopub.execute_input":"2022-05-25T08:18:37.276403Z","iopub.status.idle":"2022-05-25T08:18:37.28586Z","shell.execute_reply.started":"2022-05-25T08:18:37.276374Z","shell.execute_reply":"2022-05-25T08:18:37.28504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DolgLocalBranch(nn.Module):\n    def __init__(self, in_channel, out_channel, hidden_channel=2048):\n        super().__init__()\n        self.multi_atrous = MultiAtrous(in_channel, hidden_channel, size=int(Config.image_size/8))\n        self.conv1x1_1 = nn.Conv2d(hidden_channel, out_channel, kernel_size=1)\n        self.conv1x1_2 = nn.Conv2d(\n            out_channel, out_channel, kernel_size=1, bias=False)\n        self.conv1x1_3 = nn.Conv2d(out_channel, out_channel, kernel_size=1)\n\n        self.relu = nn.ReLU()\n        self.bn = nn.BatchNorm2d(out_channel)\n        self.softplus = nn.Softplus()\n\n    def forward(self, x):\n        local_feat = self.multi_atrous(x)\n\n        local_feat = self.conv1x1_1(local_feat)\n        local_feat = self.relu(local_feat)\n        local_feat = self.conv1x1_2(local_feat)\n        local_feat = self.bn(local_feat)\n\n        attention_map = self.relu(local_feat)\n        attention_map = self.conv1x1_3(attention_map)\n        attention_map = self.softplus(attention_map)\n\n        local_feat = F.normalize(local_feat, p=2, dim=1)\n        local_feat = local_feat * attention_map\n\n        return local_feat\n","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:18:40.001354Z","iopub.execute_input":"2022-05-25T08:18:40.001733Z","iopub.status.idle":"2022-05-25T08:18:40.01288Z","shell.execute_reply.started":"2022-05-25T08:18:40.001698Z","shell.execute_reply":"2022-05-25T08:18:40.012086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class OrthogonalFusion(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n    def forward(self, local_feat, global_feat):\n        global_feat_norm = torch.norm(global_feat, p=2, dim=1)\n        projection = torch.bmm(global_feat.unsqueeze(1), torch.flatten(\n            local_feat, start_dim=2))\n        projection = torch.bmm(global_feat.unsqueeze(\n            2), projection).view(local_feat.size())\n        projection = projection / \\\n            (global_feat_norm * global_feat_norm).view(-1, 1, 1, 1)\n        orthogonal_comp = local_feat - projection\n        global_feat = global_feat.unsqueeze(-1).unsqueeze(-1)\n        return torch.cat([global_feat.expand(orthogonal_comp.size()), orthogonal_comp], dim=1)\n","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:18:42.709243Z","iopub.execute_input":"2022-05-25T08:18:42.709539Z","iopub.status.idle":"2022-05-25T08:18:42.718197Z","shell.execute_reply.started":"2022-05-25T08:18:42.709507Z","shell.execute_reply":"2022-05-25T08:18:42.717313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    base_dir = Path('/kaggle/input/pytorch-challange-flower-dataset/')\n    cat2name_path = base_dir.joinpath('cat_to_name.json') \n    train_path = base_dir.joinpath('dataset/train')\n    val_path = base_dir.joinpath('dataset/valid')\n    test_path = base_dir.joinpath('dataset/test')\n    train_batch_size = 10\n    val_batch_size = 10\n    num_workers = 8\n    image_size = 224\n    output_dim = 102\n    hidden_dim = 1024\n    input_dim = 3\n    epochs = 35\n    lr = 1e-4\n    num_of_classes = 102\n    pretrained = True\n    model_name = 'resnet18'\n    seed = 42\n    image_transform = A.Compose([\n    A.Resize(image_size, image_size),\n    A.Normalize(),\n    A.pytorch.transforms.ToTensorV2()\n    ])\n","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:18:45.365881Z","iopub.execute_input":"2022-05-25T08:18:45.366186Z","iopub.status.idle":"2022-05-25T08:18:45.37516Z","shell.execute_reply.started":"2022-05-25T08:18:45.366146Z","shell.execute_reply":"2022-05-25T08:18:45.37356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_lightning import LightningModule\nimport torchmetrics\nimport sys\nclass DolgNet(LightningModule):\n    def __init__(self, input_dim, hidden_dim, output_dim, num_of_classes):\n        super().__init__()\n        self.cnn = timm.create_model(\n            Config.model_name,\n            pretrained=True,\n            features_only=True,\n            in_chans=input_dim,\n            out_indices=(2, 3)\n        )\n        self.cnn1 = timm.create_model(\n            Config.model_name,\n            pretrained=True,\n            in_chans=input_dim,\n        )\n        self.orthogonal_fusion = OrthogonalFusion()\n        self.local_branch = DolgLocalBranch(128, hidden_dim)\n        self.gap = nn.AdaptiveAvgPool2d(1)\n        self.gem_pool = GeM()\n        self.fc_1 = nn.Linear(256, hidden_dim)\n        self.fc_2 = nn.Linear(int(2*hidden_dim), output_dim)\n        self.fc_3 = nn.Linear(1000, output_dim)\n        self.criterion = ArcFace(\n            in_features=output_dim,\n            out_features=num_of_classes,\n            scale_factor=30,\n            margin=0.15,\n            criterion=nn.CrossEntropyLoss()\n        )\n        self.accuracy = torchmetrics.Accuracy()\n        self.lr = Config.lr\n        self.num_correct=0\n        self.num_correct_val=0\n        self.total_train_size=0\n        self.total_val_size=0\n        self.total_test_size=0\n\n    def forward(self, x):\n        #output = self.cnn(x)\n\n        #local_feat = self.local_branch(output[0])  # ,hidden_channel,16,16\n        #global_feat = self.fc_1(self.gem_pool(output[1]).squeeze())  # ,1024\n\n        #feat = self.orthogonal_fusion(local_feat, global_feat)\n        #feat = self.gap(feat).squeeze()\n        #feat = self.fc_2(feat)\n        output = self.cnn1(x)\n        feat = self.fc_3(output)\n        return feat\n\n    def training_step(self, batch, batch_idx):\n        img, label = batch\n        embd = self(img)\n        loss, logits = self.criterion(embd, label)\n        pred = logits.argmax(dim=1)\n        self.num_correct += torch.eq(pred, label).sum().float().item()\n        return {'loss': loss, 'num_correct':self.num_correct}\n    \n    def training_epoch_end(self, training_step_outputs):\n        print(self.num_correct,self.total_train_size)\n        print('train_acc',self.num_correct/self.total_train_size)\n        self.num_correct=0\n\n    def validation_step(self, batch, batch_idx):\n        img, label = batch\n        embd = self(img)\n        loss, logits = self.criterion(embd, label)\n        pred = logits.argmax(dim=1)\n        self.num_correct_val += torch.eq(pred, label).sum().float().item()\n        return {'loss': loss, 'num_correct':self.num_correct_val}\n        \n        \n    def validation_epoch_end(self, validation_step_outputs):\n        print(self.num_correct_val,self.total_val_size)\n        print('val_acc',self.num_correct_val/self.total_val_size)\n        self.num_correct_val=0\n    def test_step(self, batch, batch_idx):\n        img, label = batch\n        embd = self(img)\n        loss = F.cross_entropy(embd, label)\n        preds = embd.argmax(dim=1)\n        return {'loss': loss, 'labels': label, 'preds': preds}\n        \n    def configure_optimizers(self):\n        optimizer = optim.SGD(self.parameters(), lr=self.lr,\n                              momentum=0.9, weight_decay=1e-5)\n        scheduler = scheduler = optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=1000)\n        return [optimizer], [scheduler]\n\n    def train_dataloader(self):\n        dataset = LmkRetrDataset(Config.train_path, Config.cat2name_path, Config.image_transform)\n        train_loader=DataLoader(dataset, batch_size=Config.train_batch_size, num_workers=Config.num_workers,\n                          shuffle=True, pin_memory=True, persistent_workers=True)\n        self.total_train_size=len(train_loader.dataset)\n        return train_loader\n    def val_dataloader(self):\n        dataset = LmkRetrDataset(Config.val_path, Config.cat2name_path, Config.image_transform)\n        val_loader=DataLoader(dataset, batch_size=Config.train_batch_size,\n                          shuffle=True,num_workers=Config.num_workers,pin_memory=True, persistent_workers=True)\n        self.total_val_size=len(val_loader.dataset)\n        return val_loader\n    def test_dataloader(self):\n        dataset = LmkRetrDataset(Config.test_path, Config.cat2name_path, Config.image_transform)\n        test_loader=DataLoader(dataset, batch_size=Config.train_batch_size, num_workers=Config.num_workers,\n                          shuffle=True, pin_memory=True, persistent_workers=True)\n        self.total_test_size=len(test_loader.dataset)\n        return test_loader\n        ","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:20:21.568348Z","iopub.execute_input":"2022-05-25T08:20:21.56868Z","iopub.status.idle":"2022-05-25T08:20:21.59714Z","shell.execute_reply.started":"2022-05-25T08:20:21.568646Z","shell.execute_reply":"2022-05-25T08:20:21.596091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_lightning.utilities.seed import seed_everything\nfrom pytorch_lightning import Trainer\nseed_everything(Config.seed)\n\nmodel = DolgNet(\n    input_dim=Config.input_dim,\n    hidden_dim=Config.hidden_dim,\n    output_dim=Config.output_dim,\n    num_of_classes=Config.num_of_classes\n)\n\ntrainer = Trainer(max_epochs=Config.epochs)\n\ntrainer.fit(model)","metadata":{"execution":{"iopub.status.busy":"2022-05-25T08:20:26.211094Z","iopub.execute_input":"2022-05-25T08:20:26.211367Z"},"trusted":true},"execution_count":null,"outputs":[]}]}