{"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":"#training methods:\n#-> normal classification\n#-> arc classification\n#-> encoding\n#-> trained embedding\n#-> trained retrieval","metadata":{"execution":{"iopub.status.busy":"2022-09-19T11:13:29.277966Z","iopub.execute_input":"2022-09-19T11:13:29.278404Z","iopub.status.idle":"2022-09-19T11:13:29.283526Z","shell.execute_reply.started":"2022-09-19T11:13:29.278370Z","shell.execute_reply":"2022-09-19T11:13:29.282046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /kaggle/input/embedding-data/embeddings","metadata":{"execution":{"iopub.status.busy":"2022-09-19T11:13:29.289350Z","iopub.execute_input":"2022-09-19T11:13:29.289728Z","iopub.status.idle":"2022-09-19T11:13:30.417164Z","shell.execute_reply.started":"2022-09-19T11:13:29.289696Z","shell.execute_reply":"2022-09-19T11:13:30.415491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport cv2\nimport sys\nimport math\nimport copy\nimport torch\nimport random\nimport numpy as np\nimport pandas as pd\nfrom torch import nn\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom zipfile import ZipFile\nfrom torch.nn import LayerNorm\nimport matplotlib.pyplot as plt\nimport torch.nn.functional as F\nfrom torchvision import transforms\n\n    \ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-09-23T23:11:58.882711Z","iopub.execute_input":"2022-09-23T23:11:58.883270Z","iopub.status.idle":"2022-09-23T23:12:01.018925Z","shell.execute_reply.started":"2022-09-23T23:11:58.883157Z","shell.execute_reply":"2022-09-23T23:12:01.017789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /kaggle/input/clipembeddings","metadata":{"execution":{"iopub.status.busy":"2022-09-23T22:48:41.912560Z","iopub.execute_input":"2022-09-23T22:48:41.912988Z","iopub.status.idle":"2022-09-23T22:48:43.007942Z","shell.execute_reply.started":"2022-09-23T22:48:41.912954Z","shell.execute_reply":"2022-09-23T22:48:43.006757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.read_parquet('/kaggle/input/clipembeddings/metadata_0000.parquet', engine='pyarrow')","metadata":{"execution":{"iopub.status.busy":"2022-09-23T22:48:55.803819Z","iopub.execute_input":"2022-09-23T22:48:55.804245Z","iopub.status.idle":"2022-09-23T22:49:00.668913Z","shell.execute_reply.started":"2022-09-23T22:48:55.804206Z","shell.execute_reply":"2022-09-23T22:49:00.667988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ArcMarginProduct(nn.Module):\n    r\"\"\"Implement of large margin arc distance: :\n        Args:\n            in_features: size of each input sample\n            out_features: size of each output sample\n            s: norm of input feature\n            m: margin\n            cos(theta + m)\n        \"\"\"\n    def __init__(self, in_features, out_features, s=30.0, \n                 m=0.5, easy_margin=False, ls_eps=0.0):\n        super(ArcMarginProduct, self).__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.s = s\n        self.m = m\n        self.ls_eps = ls_eps  # label smoothing\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n\n        self.easy_margin = easy_margin\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\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        if self.easy_margin:\n            phi = torch.where(cosine > 0, phi, cosine)\n        else:\n            phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        # --------------------------- convert label to one-hot ---------------------\n        # one_hot = torch.zeros(cosine.size(), requires_grad=True, device='cuda')\n        one_hot = torch.zeros(cosine.size(), device=device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        if self.ls_eps > 0:\n            one_hot = (1 - self.ls_eps) * one_hot + self.ls_eps / self.out_features\n        # -------------torch.where(out_i = {x_i if condition_i else y_i) ------------\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-09-19T11:13:30.434625Z","iopub.execute_input":"2022-09-19T11:13:30.435441Z","iopub.status.idle":"2022-09-19T11:13:30.450200Z","shell.execute_reply.started":"2022-09-19T11:13:30.435396Z","shell.execute_reply":"2022-09-19T11:13:30.448937Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GraphAttentionLayer(nn.Module):\n\n    def __init__(self, in_features, out_features, dropout=0.6, alpha=0.2, concat=True):\n        super().__init__()\n        self.dropout = dropout\n        self.in_features = in_features\n        self.out_features = out_features\n        self.alpha = alpha\n        self.concat = concat\n\n        self.W = nn.Parameter(torch.empty(size=(in_features, out_features)))\n        nn.init.xavier_uniform_(self.W.data, gain=1.414)\n        self.a = nn.Parameter(torch.empty(size=(2 * out_features, 1)))\n        nn.init.xavier_uniform_(self.a.data, gain=1.414)\n\n        self.leakyrelu = nn.LeakyReLU(self.alpha)\n\n    def forward(self, h):\n        Wh = h @ self.W  # h.shape: (B, N, in_features), Wh.shape: (B, N, out_features)\n        a_input = self._prepare_attentional_mechanism_input(Wh)\n        e = self.leakyrelu(torch.matmul(a_input, self.a).squeeze(3))\n\n        attention = F.softmax(e, dim=1)\n        attention = F.dropout(attention, self.dropout, training=self.training)\n        h_prime = torch.bmm(attention, Wh)\n\n        if self.concat:\n            return F.elu(h_prime)\n        else:\n            return h_prime\n\n    def _prepare_attentional_mechanism_input(self, Wh):\n        B, N, D = Wh.shape\n\n        Wh_repeated_in_chunks = Wh.repeat_interleave(N, dim=1)\n        Wh_repeated_alternating = Wh.repeat(1, N, 1)\n\n        all_combinations_matrix = torch.cat([Wh_repeated_in_chunks, Wh_repeated_alternating], dim=2)\n        return all_combinations_matrix.view(-1, N, N, 2 * D)\n\n    def __repr__(self):\n        return self.__class__.__name__ + ' (' + str(self.in_features) + ' -> ' + str(self.out_features) + ')'","metadata":{"execution":{"iopub.status.busy":"2022-09-19T11:13:30.453548Z","iopub.execute_input":"2022-09-19T11:13:30.454066Z","iopub.status.idle":"2022-09-19T11:13:30.470867Z","shell.execute_reply.started":"2022-09-19T11:13:30.454019Z","shell.execute_reply":"2022-09-19T11:13:30.469430Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TopTrainer(nn.Module):\n    def __init__(self, fw):\n        super().__init__()\n        self.fw = fw\n        self.lid = None\n        self.pool = nn.AdaptiveAvgPool1d(64)\n        self.arc = None\n    def arc_sim(self, x, label):\n        x = self.fw(x)\n        return self.arc( x, label)\n    def with_arc(self, x, label, use_lid=True):\n        x = self.fw(x)\n        if use_lid:\n            x = self.lid(x)\n        x = self.arc( x, label)\n        return x\n        \n    def forward(self, x, use_lid=True):\n        x = self.fw(x) \n        if use_lid:\n            x = self.lid(x)\n        return x\n\nmodel = TopTrainer(\n    nn.Sequential(\n            nn.Linear( 768, 2048),\n            nn.Linear( 2048, 2048),\n            nn.Linear( 2048, 2048),\n            nn.Linear( 2048, 2048),\n            nn.Linear( 2048, 2048),\n            nn.Linear( 2048, 64)\n        )\n).to(device)\nmodel","metadata":{"execution":{"iopub.status.busy":"2022-09-23T23:15:47.436807Z","iopub.execute_input":"2022-09-23T23:15:47.437464Z","iopub.status.idle":"2022-09-23T23:15:47.618013Z","shell.execute_reply.started":"2022-09-23T23:15:47.437426Z","shell.execute_reply":"2022-09-23T23:15:47.616819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class class_ds_config:\n    def __init__(\n            self,\n            path,\n            lr,\n            epochs,\n            train, \n            max_class\n    ):\n        self.path=path\n        self.lr=lr\n        self.epochs=epochs\n        self.train=train\n        self.max_class = max_class\nclass embed_ds_config:\n    def __init__(\n        self,\n        path,\n        lr,\n        epochs,\n        train,\n        criterion\n    ):\n        self.path=path\n        self.lr=lr\n        self.epochs=epochs\n        self.train=train,\n        self.criterion=criterion\nclass config:\n    \n    class_data_ls = [     \n        class_ds_config(\n            '/kaggle/input/embedding-data/embeddings/architecture-dataset',\n            0.0001, 2, True, 25\n        ),\n        class_ds_config(\n            '/kaggle/input/notebook-data/classification/GPR',\n            0.00001,2,True,1200\n        ),\n        class_ds_config(\n            '/kaggle/input/embedding-data/embeddings/food101',\n            0.00008, 2, True, 101\n        ),\n        class_ds_config(\n            '/kaggle/input/embedding-data/embeddings/fruits360',\n            0.0001, 1, True, 131\n        ),\n        class_ds_config(\n            '/kaggle/input/embedding-data/embeddings/rp2k',\n            0.0001, 15, False, 2384\n        ),\n        class_ds_config(\n            '/kaggle/input/embedding-data/embeddings/products-10k',\n            0.0001, 5, True, 9691\n        ),\n        class_ds_config(\n            '/kaggle/input/embedding-data/embeddings/shopee',\n            0.0001, 10, True, 11014\n        ),\n        class_ds_config(\n            '/kaggle/input/embedding-data/embeddings/imagenet1000',\n            0.0001, 10, True, 1000\n        ),\n        class_ds_config(\n            '/kaggle/input/embedding-data/embeddings/places',\n            0.0001, 2, True, 1000\n        ),\n        \n        class_ds_config(\n            '/kaggle/input/embedding-data/embeddings/imagenet-sketch',\n            0.0001, 2, True, 1000\n        ),\n        class_ds_config(\n            '/kaggle/input/notebook-data/classification/imagenet1000',\n            0.0001,3,False,1000\n        )\n    ]\n    embed_data_ls = [\n        embed_ds_config(\n            '/kaggle/input/notebook-data/embedding/imagenet1000',\n            0.0005,\n            3,\n            False,\n            nn.MSELoss()\n        ),\n        embed_ds_config(\n            '/kaggle/input/notebook-data/embedding/google-landmarks-2021-V1',\n            0.0005,\n            2,\n            False,\n            nn.MSELoss()\n        ),\n        embed_ds_config(\n            '/kaggle/input/notebook-data/embedding/fashion',\n            0.00005,\n            2,\n            False,\n            nn.MSELoss()\n        ),\n        embed_ds_config(\n            '/kaggle/input/notebook-data/embedding/GPR12000',\n            0.0005,\n            2,\n            False,\n            nn.MSELoss()\n        )\n    ]","metadata":{"execution":{"iopub.status.busy":"2022-09-23T23:12:02.940085Z","iopub.execute_input":"2022-09-23T23:12:02.940539Z","iopub.status.idle":"2022-09-23T23:12:02.954535Z","shell.execute_reply.started":"2022-09-23T23:12:02.940499Z","shell.execute_reply":"2022-09-23T23:12:02.953356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def retrieval_evaluate( model, ds, ds_config):\n    def get_embeds( model, ds):\n        embeds = []\n        labels = []\n        with torch.no_grad():\n            for vec, labels_ in ds:\n                out = model( vec.to(device), use_lid=False)\n                embeds.append(out)\n                labels.append(labels_)\n        embeds = torch.cat(embeds)\n        labels = torch.cat(labels)\n        return (embeds, labels)\n    def normalize(a, eps=1e-8):\n        a_n = a.norm(dim=1)[:, None]\n        a_norm = a / torch.max(a_n, eps * torch.ones_like(a_n))\n        return a_norm\n    def k_nearest_neighbors(embeds, k=5):\n        #print(embeds.shape)\n        normalized = normalize(embeds)\n\n        preds = normalized @ normalized.T\n        #vals, indices = preds.sort(dim=1, descending=True)\n        vals,indices = torch.topk(preds,6)\n        k += 1\n        return indices[:, 1:k].long()\n    embeds, labels = get_embeds( model, ds)\n    preds = k_nearest_neighbors(embeds, k=5)\n    accs = (labels[preds] == labels.view(-1, 1)).float().mean(dim=1)\n    return accs.mean()\n\nlast_retrieval = [0 for i in config.class_data_ls]\ndef get_retrieval_scores():\n    idx=0\n    for ds_config in config.class_data_ls:\n        if ds_config.train:\n            ds = classification_dataset(ds_config)\n            if ds_config.path.split('/')[-1] == 'products-10k':\n                #print('splicing ds...',end='    ')\n                ds = ds.splice(0,512)\n            #ds = ds.splice(0,512)\n            #val_score = retrieval_evaluate( model, ds.val_split(), ds_config)\n            #train_score = retrieval_evaluate( model, ds.train_split(), ds_config)\n            score = retrieval_evaluate( model, ds, ds_config)\n            diff = score.item() - last_retrieval[idx]\n            delta = '+' if diff >= 0 else ' ' \n            pad = 20 - len(ds_config.path.split('/')[-1])\n            print( ds_config.path.split('/')[-1], \" \"*pad,\": \", score.item(),\"  {}{}\".format(delta, diff))\n            last_retrieval[idx] = score.item()\n            #print( ds_config.path.split('/')[-1],\" train: \", train_score.item())\n            #print( ds_config.path.split('/')[-1],\"   val: \", val_score.item())\n            idx+=1\ndef _eval( model):\n    for ds_config in config.class_data_ls[:1]:\n        ds = classification_dataset(ds_config)\n        ds = ds.val_split()\n        score = retrieval_evaluate( model, ds, ds_config)\n        print( ds_config.path.split('/')[-1],\": \", score.item())","metadata":{"execution":{"iopub.status.busy":"2022-09-23T23:12:03.927826Z","iopub.execute_input":"2022-09-23T23:12:03.928987Z","iopub.status.idle":"2022-09-23T23:12:03.948359Z","shell.execute_reply.started":"2022-09-23T23:12:03.928940Z","shell.execute_reply":"2022-09-23T23:12:03.947004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class classification_dataset(torch.utils.data.Dataset):\n    def __init__(self, config):\n        self.config = config\n        self.vecs = ['vec/'+_dir for _dir in os.listdir(os.path.join(config.path, 'vec'))]\n        self.vecs.sort()\n        self.labels = ['label/'+_dir for _dir in os.listdir(os.path.join(config.path, 'label'))]\n        self.labels.sort()\n    def __len__(self):\n        return len(self.vecs)\n    def train_split(self):\n        _len = len(self)\n        split = int(_len * 0.8)\n        self.vecs = self.vecs[:split]\n        self.labels = self.labels[:split]\n        return self\n    def val_split(self):\n        _len = len(self)\n        split = int(_len * 0.8)\n        self.vecs = self.vecs[split:]\n        self.labels = self.labels[split:]\n        return self\n        \n    def splice(self, start, end):\n        self.vecs = self.vecs[start:end]\n        self.labels = self.labels[start:end]\n        return self\n    def __getitem__(self, idx):\n        if self.vecs[idx].split('/')[1] != self.labels[idx].split('/')[1]:\n            print('error')\n        vec = torch.load(os.path.join(self.config.path,self.vecs[idx]),map_location=device)\n        label = torch.load(os.path.join(self.config.path,self.labels[idx]),map_location=device)\n        return vec, label","metadata":{"execution":{"iopub.status.busy":"2022-09-23T23:12:11.449006Z","iopub.execute_input":"2022-09-23T23:12:11.449408Z","iopub.status.idle":"2022-09-23T23:12:11.464107Z","shell.execute_reply.started":"2022-09-23T23:12:11.449374Z","shell.execute_reply":"2022-09-23T23:12:11.462929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"get_retrieval_scores()","metadata":{"execution":{"iopub.status.busy":"2022-09-19T11:13:30.624339Z","iopub.execute_input":"2022-09-19T11:13:30.624737Z"},"jupyter":{"source_hidden":true,"outputs_hidden":true},"collapsed":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SAM(torch.optim.Optimizer):\n    def __init__(self, params, base_optimizer, rho=0.05, **kwargs):\n        assert rho >= 0.0, f\"Invalid rho, should be non-negative: {rho}\"\n\n        defaults = dict(rho=rho, **kwargs)\n        super(SAM, self).__init__(params, defaults)\n\n        self.base_optimizer = base_optimizer(self.param_groups, **kwargs)\n        self.param_groups = self.base_optimizer.param_groups\n\n    @torch.no_grad()\n    def first_step(self, zero_grad=False):\n        grad_norm = self._grad_norm()\n        for group in self.param_groups:\n            scale = group[\"rho\"] / (grad_norm + 1e-12)\n\n            for p in group[\"params\"]:\n                if p.grad is None: continue\n                e_w = p.grad * scale.to(p)\n                p.add_(e_w)  # climb to the local maximum \"w + e(w)\"\n                self.state[p][\"e_w\"] = e_w\n\n        if zero_grad: self.zero_grad()\n\n    @torch.no_grad()\n    def second_step(self, zero_grad=False):\n        for group in self.param_groups:\n            for p in group[\"params\"]:\n                if p.grad is None: continue\n                p.sub_(self.state[p][\"e_w\"])  # get back to \"w\" from \"w + e(w)\"\n\n        self.base_optimizer.step()  # do the actual \"sharpness-aware\" update\n\n        if zero_grad: self.zero_grad()\n\n    def step(self, closure=None):\n        raise NotImplementedError(\"SAM doesn't work like the other optimizers, you should first call `first_step` and the `second_step`; see the documentation for more info.\")\n\n    def _grad_norm(self):\n        shared_device = self.param_groups[0][\"params\"][0].device  # put everything on the same device, in case of model parallelism\n        norm = torch.norm(\n                    torch.stack([\n                        p.grad.norm(p=2).to(shared_device)\n                        for group in self.param_groups for p in group[\"params\"]\n                        if p.grad is not None\n                    ]),\n                    p=2\n               )\n        return norm","metadata":{"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_arc_ds( model, ds, ds_config):    \n    \n    optim = torch.optim.Adam(model.parameters(), lr=ds_config.lr)\n    #base_optimizer = torch.optim.Adam\n    #optimizer = SAM(model.parameters(), base_optimizer, lr=0.001)  \n    criterion = nn.CrossEntropyLoss()\n    model.train()\n    \n    for epoch in range(ds_config.epochs):\n        print('training epoch ',epoch+1,'...',end='')\n        total_loss = 0\n        for i, batch in enumerate(ds):\n            _input, label = batch\n            \n            #image_preds = model.with_arc(_input.to(device),label.to(device))   #output = model(input)\n            #print(image_preds.shape, exam_pred.shape)\n\n            #oss = criterion(image_preds, label.to(device)) \n            #loss.backward()\n            #optimizer.first_step(zero_grad=True)\n\n            # second forward-backward pass\n            #criterion(model.with_arc(_input.to(device), label.to(device)), label.to(device)).backward()\n            #optimizer.second_step(zero_grad=True)\n            #total_loss += loss.detach().item()\n\n            _input = _input.to(device)\n            optim.zero_grad()\n            output = model.with_arc( _input, label)\n\n            loss = criterion( output, label.to(device))\n            loss.backward()\n            optim.step()\n        print('      {}'.format(total_loss))\n        total_loss = 0\n    return model\nif False:\n    get_retrieval_scores()\n    for ds_config in config.class_data_ls:\n        if ds_config.train:\n            print('='*50)\n            ds = classification_dataset(ds_config)\n            ds = ds.train_split()\n            model.lid = nn.Linear(64, ds_config.max_class).to(device)\n            model.arc = ArcMarginProduct( ds_config.max_class,ds_config.max_class).to(device)\n            \n\n            print('training dataset: {}'.format(ds_config.path.split('/')[-1]))\n            model = train_arc_ds( model, ds, ds_config)\n            \n            #get_retrieval_scores()\n            print('finished.')\n            print('='*50)\n            print('')\n    get_retrieval_scores()","metadata":{"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2022-09-23T23:35:24.483965Z","iopub.status.idle":"2022-09-23T23:35:24.484412Z","shell.execute_reply.started":"2022-09-23T23:35:24.484205Z","shell.execute_reply":"2022-09-23T23:35:24.484226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2022-09-23T22:54:02.445316Z","iopub.execute_input":"2022-09-23T22:54:02.445818Z","iopub.status.idle":"2022-09-23T22:54:02.476646Z","shell.execute_reply.started":"2022-09-23T22:54:02.445781Z","shell.execute_reply":"2022-09-23T22:54:02.475545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_encoding_ds(model, ds, ds_config):\n    optim = torch.optim.SGD(model.fw.parameters(), lr=ds_config.lr)\n    criterion = nn.MSELoss()#ds_config.criterion\n    model.train()\n    total_loss=0\n    for i, batch in enumerate(ds):\n        optim.zero_grad()\n        model.zero_grad()\n        output = model(batch.to(device))\n        loss = criterion( output, batch) \n        loss.backward()\n        optim.step()\n        total_loss += loss.detach().item()\n    print('     {}'.format(total_loss / len(ds)))\n    return model\nget_retrieval_scores()\nfor ds_config in config.embed_data_ls:\n    if ds_config.train:\n        model.lid = nn.Linear(64,768).to(device)\n        ds = embedding_dataset(ds_config)\n        print('training ',ds_config.path)\n        \n        for epoch in range(ds_config.epochs):\n            print('{} / {} training {}...'.format( epoch+1, ds_config.epochs, ds_config.path),end='')\n            model = train_encoding_ds( model, ds, ds_config)\n        \n        print('')\nget_retrieval_scores()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class JSD(nn.Module):\n    #thanks to:\n    #https://discuss.pytorch.org/t/jensen-shannon-divergence/2626/12\n    def __init__(self):\n        super(JSD, self).__init__()\n        self.kl = nn.KLDivLoss(reduction='batchmean', log_target=True)\n\n    def forward(self, p: torch.tensor, q: torch.tensor):\n        p, q = p.view(-1, p.size(-1)).log_softmax(-1), q.view(-1, q.size(-1)).log_softmax(-1)\n        m = (0.5 * (p + q))\n        return 0.5 * (self.kl(m, p) + self.kl(m, q))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_encoding_ds(model):\n    vecs = np.load('/kaggle/input/clipembeddings/img_emb_0000.npy')\n    ds = torch.utils.data.DataLoader(\n        vecs,\n        batch_size = 128\n    )\n    \n    optim = torch.optim.Adam(model.parameters(), lr=0.0001)\n    criterion = nn.MSELoss()\n    model.train()\n    total_loss=0\n    for epoch in range(1):\n        for i, batch in enumerate(ds):\n            optim.zero_grad()\n            model.zero_grad()\n\n            output = model(batch.to(device).float())\n            loss = criterion( output, batch.to(device).float()) \n            loss.backward()\n            optim.step()\n            total_loss += loss.detach().item()\n            print('-',end='')\n            if i % 150 == 0 and i is not 0:\n                print('{} / {}, last step loss: {}'.format(i, len(ds), loss.item()))\n        print('     {}'.format(total_loss / len(ds)))\n    return model\n#get_retrieval_scores()\nmodel.lid = nn.Linear(64,768).to(device)\nmodel = train_encoding_ds( model)\nget_retrieval_scores()","metadata":{"execution":{"iopub.status.busy":"2022-09-23T23:35:27.223142Z","iopub.execute_input":"2022-09-23T23:35:27.223551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class classification_dataset(torch.utils.data.Dataset):\n    def __init__(self, config):\n        self.config = config\n        self.vecs = ['vec/'+_dir for _dir in os.listdir(os.path.join(config.path, 'vec'))]\n        self.vecs.sort()\n        self.labels = ['label/'+_dir for _dir in os.listdir(os.path.join(config.path, 'label'))]\n        self.labels.sort()\n    def __len__(self):\n        return len(self.vecs)\n    def train_split(self):\n        _len = len(self)\n        split = int(_len * 0.8)\n        self.vecs = self.vecs[:split]\n        self.labels = self.labels[:split]\n        return self\n    def val_split(self):\n        _len = len(self)\n        split = int(_len * 0.8)\n        self.vecs = self.vecs[split:]\n        self.labels = self.labels[split:]\n        return self\n        \n    def splice(self, start, end):\n        self.vecs = self.vecs[start:end]\n        self.labels = self.labels[start:end]\n        return self\n    def __getitem__(self, idx):\n        if self.vecs[idx].split('/')[1] != self.labels[idx].split('/')[1]:\n            print('error')\n        vec = torch.load(os.path.join(self.config.path,self.vecs[idx]),map_location=device)\n        label = torch.load(os.path.join(self.config.path,self.labels[idx]),map_location=device)\n        return vec, label","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CLIPEmbeddingDS(torch.utils.data.Dataset):\n    def __init__(self):\n        \n        self.text = pd.read_parquet('/kaggle/input/clipembeddings/metadata_0000.parquet', engine='pyarrow').numpy()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.read_parquet('/kaggle/input/clipembeddings/metadata_0000.parquet', engine='pyarrow')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install ftfy regex tqdm\n!pip install git+https://github.com/openai/CLIP.git\n\n\nimport torch\nimport clip\nfrom clip.clip import _download, _MODELS\n\n\nmodel_path = _download(_MODELS['ViT-L/14@336px'], os.path.expanduser(\"~/.cache/clip\"))\nwith open(model_path, 'rb') as opened_file:\n    print('opening: ',model_path)\n    clip_vit_l14_336 = torch.jit.load(opened_file, map_location=device).visual.eval()\nclass MyModel(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = clip_vit_l14_336\n        \n        self.fw = None\n        self.pool = nn.AdaptiveAvgPool1d(64)\n    def preprocess_image(self, x):\n        x = transforms.functional.resize(x,size=[336, 336])\n        x = x/255.0\n        x = transforms.functional.normalize(x, \n                                            mean=[0.48145466, 0.4578275, 0.40821073],\n                                            std=[0.26862954, 0.26130258, 0.27577711])\n        return x\n    \n    def forward(self, x):\n        x = self.preprocess_image(x)\n        x = self.encoder(x.half())\n        x = self.fw(x)\n        #x = torch.nn.functional.normalize(x, p=2.0, dim=1, eps=1e-12)\n        return x\n\nsub = MyModel().to(device).eval()\nsub.fw = model.fw.half()\n#print(sub)\nprint(sub(torch.randn(1,3,336,336).to(device)).detach())\n\nsub.eval()\nsaved_model = torch.jit.script(sub)\nsaved_model.save(\"saved_model.pt\")\nwith ZipFile('submission.zip','w') as zip:           \n    zip.write(\"saved_model.pt\", arcname='saved_model.pt')\nsub = torch.jit.load(\"saved_model.pt\").to('cuda').eval()\ninput_batch = torch.rand(1, 3, 336, 336).to('cuda')\nwith torch.no_grad():\n    embedding = sub(input_batch).cpu().data.numpy()\nembedding","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.fw","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if False:#config.train_embed:\n    \n    #train on imagenet classes\n    #path = '/kaggle/input/imagenetmini-1000/imagenet-mini/train'\n    #classes = os.listdir(path)\n    #idx=0\n    #for _class in classes[:config.imagenet_classes]:\n    #    idx+=1\n    #    print(' imagenetmini1000 | trainig class {} | {} / {} ...'.format( _class, idx, config.imagenet_classes))\n    #    path_ds = load_imagenet_class(os.path.join(path, _class))\n    #    train_class( path_ds, model)\n    idx=0 \n    for _class in get_imagenet1000_classes()[:config.imagenet_embed_classes]:\n        idx+=1\n        print(' imagenet1000 | training embed class {} | {} / {} ...'.format( _class, idx, 200))\n        path_ds = load_imagenet1000_class(_class)\n        train_embed_ds( path_ds, model)\n    #train on caltech256 classes\n    idx=0\n    for _class in get_caltech256_classes()[:config.caltech256_embed_classes]:\n        idx+=1\n        print(' caltech256 | training embed class {} | {} / {} ...'.format( _class, idx, 20))\n        path_ds = load_caltech256_class(_class)\n        train_embed_ds( path_ds, model)\n    idx=0\n    for _class in get_fashion_classes()[:config.fashion_embed_classes]:\n        idx+=1\n        print(' caltech256 | training embed class {} | {} / {} ...'.format( _class, idx, config.fasion_embed_classes))\n        path_ds = load_fashion_class(_class)\n        train_embed_ds( path_ds, model)\n        \n        \n    \n    #train on google landmark recognition 2021 classes\n    #idx=0\n    #for _class in get_google_landmarks_2021_classes()[:500]:\n    #    idx+=1\n    #    print(' Google Landmarks 2021 | training class {} | {} / {}'.format( _class, idx, 500))\n    #    path_ds = load_google_landmarks_2021_class(_class)\n    #    train_class( path_ds, model)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_encoding_class_ds(model, ds, ds_config):\n    optim = torch.optim.SGD(model.fw.parameters(), lr=ds_config.lr)\n    criterion = nn.MSELoss()#ds_config.criterion\n    model.train()\n    total_loss=0\n    for i, batch in enumerate(ds):\n        _input, label = batch\n        optim.zero_grad()\n        model.zero_grad()\n        output = model(_input.to(device))\n        loss = criterion( output, _input.to(device)) \n        loss.backward()\n        optim.step()\n        total_loss += loss.detach().item()\n    print('     {}'.format(total_loss / len(ds)))\n    return model\n\nif False:\n    get_retrieval_scores()\n    model.lid = nn.Linear(64,768).to(device)\n    for epoch in range(10):\n        for ds_config in config.class_data_ls:\n            if ds_config.train:\n                ds = classification_dataset(ds_config)\n                print('training ',ds_config.path)\n                print('epoch {} training {}...'.format( epoch+1, ds_config.path.split('/')[-1]),end='')\n                model = train_encoding_class_ds( model, ds, ds_config)\n\n                print('')\n    get_retrieval_scores()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_classification_ds( model, ds, ds_config):\n    optim = torch.optim.Adam(model.parameters(), lr=ds_config.lr)\n    criterion = nn.CrossEntropyLoss()\n    model.train()\n    \n    for epoch in range(ds_config.epochs):\n        print('training epoch ',epoch+1,'...')\n        for i, batch in enumerate(ds):\n            _input, label = batch\n            \n            _input = _input.to(device)\n            optim.zero_grad()\n            output = model(_input)\n            \n            loss = criterion( output, label.to(device))\n            loss.backward()\n            optim.step()\n    return model\n\ndef eval_classification_ds( model, ds, ds_config):\n    criterion = nn.CrossEntropyLoss()\n    model.eval()\n    total_loss=0\n    for i, batch in enumerate(ds):\n        _input, label = batch\n        _input = _input.to(device)\n        output = model(_input)\n        loss = criterion( output, label.to(device)).detach().item()\n        total_loss += loss\n    return total_loss / len(ds)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp = model\ndef train_lookup_embed_ds( model, ds, ds_config, vec_lib):\n    optim = torch.optim.Adam(model.parameters(), lr=ds_config.lr)\n    criterion = JSD()\n    model.train()\n    \n    for epoch in range(5):#ds_config.epochs):\n        print('training epoch ',epoch+1,'...')\n        for i, batch in enumerate(ds):\n            _input, label = batch\n            \n            _input = _input.to(device)\n            optim.zero_grad()\n            model.zero_grad()\n            output = model(_input,use_lid=False)\n            \n            label = vec_lib[label]\n            loss = criterion( output, label.detach().to(device))\n            loss.backward()\n            optim.step()\n    return model\n#for ds_config in config.class_data_ls:\n#    if ds_config.train:\n#        ds = classification_dataset(ds_config)\n#        print('training ',ds_config.path)\n#        temp = train_lookup_embed_ds( temp, ds, ds_config, sample_vec_lib)\n#        print('finished.')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LookupModel(nn.Module):\n    def __init__(self, _model):\n        super().__init__()\n        self._model = _model\n        self.vec_lib = None\n    def forward( self, x, use_lid=False):\n        output = self._model(x,use_lid=False)[:,None,:]\n        output = output.repeat(1,self.vec_lib.shape[0],1)\n        vec_lib = self.vec_lib.repeat(output.shape[0],1,1)\n        product = output @ vec_lib.transpose(1,2)\n        product = product[:,:1,:]\n        value, index = torch.topk(product.squeeze(), k=1,dim=1)\n        return self.vec_lib[index.squeeze()]\n#m = LookupModel(model)\n#m.vec_lib = sample_vec_lib\n#m(torch.zeros(3,768)).shape","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_embedding_lib( model, ds, ds_config):    \n    #1. find each class in the dataset\n    all_labels = None\n    for i, batch in enumerate(ds):\n        _, label = batch\n        if all_labels is None:\n            all_labels = label\n        else:\n            all_labels = torch.concat((all_labels, label),0)\n    unique = torch.unique(all_labels)\n    \n    #2. for each class, gather all of the vectors for that class\n    vecs = None\n    for label in unique:\n        print('-',end='')\n        class_vecs = None\n        for i, batch in enumerate(ds):\n            vec, labels = batch\n            mask = (labels == label).nonzero()#torch.where(labels==label, label, -1)\n            selected = vec[mask]\n            if selected.shape[0] is not 0:\n                \n                if class_vecs is None:\n                    class_vecs = selected\n                else:\n                    class_vecs = torch.concat((class_vecs, selected),0)    \n                    \n        #3. for each set of vectors run the model on each vector, and get the average vector\n        class_vecs = class_vecs.squeeze()\n        loader = torch.utils.data.DataLoader(\n            class_vecs,\n            batch_size=64\n        )\n        output_vecs = None\n        for i, batch in enumerate(loader):\n            output = model(batch, use_lid=False)\n            if output_vecs is None:\n                output_vecs = output\n            else:\n                output_vecs = torch.concat((output_vecs, output),0)\n                \n        #4. calculate an average output vector for each class\n        mean_vec = torch.sum( output_vecs, 0) / class_vecs.shape[0]\n        mean_vec = mean_vec[None, :]\n        if vecs is None:\n            vecs = mean_vec\n        else:\n            vecs = torch.concat(( vecs, mean_vec), 0)\n    print(vecs.shape)\n    return vecs\n        \n        \n    \n        \n#for ds_config in config.class_data_ls[:1]:\n#    if ds_config.train:\n#        ds = classification_dataset(ds_config)\n#        sample_vec_lib = make_embedding_lib( model, ds, ds_config)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_embed_ds(path_ds, model):\n    \n    loader = torch.utils.data.DataLoader(\n        path_ds,\n        batch_size=config.batch_size,\n    )\n    \n    mse = torch.nn.MSELoss()\n    \n    #optim = torch.optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr=0.000005, amsgrad=True)\n    optim = torch.optim.SGD(model.fw.parameters(), lr=config.embed_lr)\n    \n    for i, batch in enumerate(loader):\n        optim.zero_grad()\n        model.zero_grad()\n        model.train()\n        output = model(batch.to(device))\n        #print('ran model')\n        \n        \n        label_vec = torch.sum( output, 0) / batch.shape[0] # calculate the average output vector\n        # essentially, make the outputs of the model more similar to each other for each class\n        # 'tighen' the vector output for the distribution of samples\n        loss = mse( output, label_vec[None,:].repeat(batch.shape[0], 1)) * 10 \n        print(loss.item())\n\n        \n        loss.backward()\n        #print(model.fw[0].weight)\n        #torch.nn.utils.clip_grad_norm_(model.parameters(), 0.5) # I have no idea what this does\n        optim.step()\n        #print(model.fw[0].weight)\n\n\nif False:#config.train_class:\n    \n    top_fw = nn.Sequential(\n            nn.Linear(768, 768),\n            nn.ReLU(),\n            nn.Linear(768, 768),\n            nn.ReLU(),\n            nn.Linear(768, 256),\n            nn.ReLU(),\n            nn.Linear(256, 64),\n            nn.ReLU(),\n            nn.Linear(64, 256),\n    ).to(device)\n    ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_embedding_ds( model, ds, ds_config):\n    optim = torch.optim.SGD(model.fw.parameters(), lr=ds_config.lr)\n    criterion = ds_config.criterion\n    model.train()\n    for i, batch in enumerate(ds):\n        optim.zero_grad()\n        model.zero_grad()\n\n        output = model(batch.to(device), use_lid=False)\n        label_vec = torch.sum( output, 0) / batch.shape[0] \n\n        loss = criterion( output, label_vec[None,:].repeat(batch.shape[0], 1))# * 10 \n\n        loss.backward()\n        optim.step()\n    return model\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class embedding_dataset(torch.utils.data.Dataset):\n    def __init__(self, config):\n        self.config = config\n        self.vecs = ['vec/'+_dir for _dir in os.listdir(os.path.join(config.path, 'vec'))]\n        self.vecs.sort()\n       \n    def __len__(self):\n        return len(self.vecs)\n\n    def __getitem__(self, idx):\n        vec = torch.load(os.path.join(self.config.path,self.vecs[idx]),map_location=device)\n        return vec\n\ndef train_embedding_ds( model, ds, ds_config):\n    optim = torch.optim.SGD(model.fw.parameters(), lr=ds_config.lr)\n    criterion = ds_config.criterion\n    model.train()\n    for i, batch in enumerate(ds):\n        optim.zero_grad()\n        model.zero_grad()\n\n        output = model(batch.to(device), use_lid=False)\n        label_vec = torch.sum( output, 0) / batch.shape[0] \n\n        loss = criterion( output, label_vec[None,:].repeat(batch.shape[0], 1))# * 10 \n\n        loss.backward()\n        optim.step()\n    return model\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#idea, have a special arcface layer to go on top of the flattened covariance matrix,\n#it will just have class 0 and 1\n#input: 64^2, output: 64^2\n\ndef train_similarities( model, ds, ds_config):\n    optim = torch.optim.Adam(model.parameters(), lr=0.0001)#ds_config.lr)\n    criterion = APLoss()\n    criterion.to(device)\n    model.train()\n    \n    for epoch in range(20):\n        print('training epoch ',epoch+1,'...',end='  ')\n        total_loss = 0\n        input_batch, label_batch = None, None\n        for i, batch in enumerate(ds):\n            _input, label = batch\n            if input_batch is None:\n                input_batch = _input\n                label_batch = label\n            else:\n                input_batch = torch.concat((input_batch, _input),0)\n                label_batch = torch.concat((label_batch, label),0)\n\n            if i%12==0 and i!=0:\n                _input = input_batch\n                label = label_batch\n                _input = _input.to(device)\n                label = label.to(device)\n                optim.zero_grad()\n                output = model( _input, use_lid=False)\n                \n                x, y = torch.meshgrid(label,label)\n                mesh = (x==y).type(torch.uint8).type(torch.float)\n                print(mesh, sum(mesh))\n                covariance = output @ output.T\n                vals,indices = torch.topk(covariance,6)[:,1:sum(mesh)]\n                #print(mesh.shape,end='  ')\n                #loss = criterion( covariance, mesh)\n                total_loss += loss.detach().item()\n\n                #loss = criterion( output, label.to(device))\n                loss.backward()\n                optim.step()\n                input_batch, label_batch = None, None\n        print('   loss: {}'.format(total_loss / len(ds)))\n    return model\nif True:\n    print('training similarities... ')\n    #get_retrieval_scores()\n    model.cuda()\n    \n    for ds_config in config.class_data_ls:\n        if ds_config.train:\n            ds = classification_dataset(ds_config)\n            print('training ',ds_config.path)\n            \n            model = train_similarities( model, ds, ds_config)\n            print('finished.')\n    #get_retrieval_scores()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class APLoss (nn.Module):\n    \"\"\" Differentiable AP loss, through quantization. From the paper:\n        Learning with Average Precision: Training Image Retrieval with a Listwise Loss\n        Jerome Revaud, Jon Almazan, Rafael Sampaio de Rezende, Cesar de Souza\n        https://arxiv.org/abs/1906.07589\n        Input: (N, M)   values in [min, max]\n        label: (N, M)   values in {0, 1}\n        Returns: 1 - mAP (mean AP for each n in {1..N})\n                 Note: typically, this is what you wanna minimize\n    \"\"\"\n    def __init__(self, nq=25, min=0, max=1):\n        nn.Module.__init__(self)\n        assert isinstance(nq, int) and 2 <= nq <= 100\n        self.nq = nq\n        self.min = min\n        self.max = max\n        gap = max - min\n        assert gap > 0\n        # Initialize quantizer as non-trainable convolution\n        self.quantizer = q = nn.Conv1d(1, 2*nq, kernel_size=1, bias=True)\n        q.weight = nn.Parameter(q.weight.detach(), requires_grad=False)\n        q.bias = nn.Parameter(q.bias.detach(), requires_grad=False)\n        a = (nq-1) / gap\n        # First half equal to lines passing to (min+x,1) and (min+x+1/a,0) with x = {nq-1..0}*gap/(nq-1)\n        q.weight[:nq] = -a\n        q.bias[:nq] = torch.from_numpy(a*min + np.arange(nq, 0, -1))  # b = 1 + a*(min+x)\n        # First half equal to lines passing to (min+x,1) and (min+x-1/a,0) with x = {nq-1..0}*gap/(nq-1)\n        q.weight[nq:] = a\n        q.bias[nq:] = torch.from_numpy(np.arange(2-nq, 2, 1) - a*min)  # b = 1 - a*(min+x)\n        # First and last one as a horizontal straight line\n        q.weight[0] = q.weight[-1] = 0\n        q.bias[0] = q.bias[-1] = 1\n\n    def forward(self, x, label, qw=None, ret='1-mAP'):\n        assert x.shape == label.shape  # N x M\n        N, M = x.shape\n        # Quantize all predictions\n        q = self.quantizer(x.unsqueeze(1))\n        q = torch.min(q[:, :self.nq], q[:, self.nq:]).clamp(min=0)  # N x Q x M\n\n        nbs = q.sum(dim=-1)  # number of samples  N x Q = c\n        rec = (q * label.view(N, 1, M).float()).sum(dim=-1)  # number of correct samples = c+ N x Q\n        prec = rec.cumsum(dim=-1) / (1e-16 + nbs.cumsum(dim=-1))  # precision\n        rec /= rec.sum(dim=-1).unsqueeze(1)  # norm in [0,1]\n\n        ap = (prec * rec).sum(dim=-1)  # per-image AP\n\n        if ret == '1-mAP':\n            if qw is not None:\n                ap *= qw  # query weights\n            return 1 - ap.mean()\n        elif ret == 'AP':\n            assert qw is None\n            return ap\n        else:\n            raise ValueError(\"Bad return type for APLoss(): %s\" % str(ret))\n\n    def measures(self, x, gt, loss=None):\n        if loss is None:\n            loss = self.forward(x, gt)\n        return {'loss_ap': float(loss)}\nclass TAPLoss (APLoss):\n    \"\"\" Differentiable tie-aware AP loss, through quantization. From the paper:\n        Learning with Average Precision: Training Image Retrieval with a Listwise Loss\n        Jerome Revaud, Jon Almazan, Rafael Sampaio de Rezende, Cesar de Souza\n        https://arxiv.org/abs/1906.07589\n        Input: (N, M)   values in [min, max]\n        label: (N, M)   values in {0, 1}\n        Returns: 1 - mAP (mean AP for each n in {1..N})\n                 Note: typically, this is what you wanna minimize\n    \"\"\"\n    def __init__(self, nq=25, min=0, max=1, simplified=False):\n        APLoss.__init__(self, nq=nq, min=min, max=max)\n        self.simplified = simplified\n\n    def forward(self, x, label, qw=None, ret='1-mAP'):\n        '''N: number of images;\n           M: size of the descs;\n           Q: number of bins (nq);\n        '''\n        assert x.shape == label.shape  # N x M\n        N, M = x.shape\n        label = label.float()\n        Np = label.sum(dim=-1, keepdim=True)\n\n        # Quantize all predictions\n        q = self.quantizer(x.unsqueeze(1))\n        q = torch.min(q[:, :self.nq], q[:, self.nq:]).clamp(min=0)  # N x Q x M\n\n        c = q.sum(dim=-1)  # number of samples  N x Q = nbs on APLoss\n        cp = (q * label.view(N, 1, M)).sum(dim=-1)  # N x Q number of correct samples = rec on APLoss\n        C = c.cumsum(dim=-1)\n        Cp = cp.cumsum(dim=-1)\n\n        zeros = torch.zeros(N, 1).to(x.device)\n        C_1d = torch.cat((zeros, C[:, :-1]), dim=-1)\n        Cp_1d = torch.cat((zeros, Cp[:, :-1]), dim=-1)\n\n        if self.simplified:\n            aps = cp * (Cp_1d+Cp+1) / (C_1d+C+1) / Np\n        else:\n            eps = 1e-8\n            ratio = (cp - 1).clamp(min=0) / ((c-1).clamp(min=0) + eps)\n            aps = cp * (c * ratio + (Cp_1d + 1 - ratio * (C_1d + 1)) * torch.log((C + 1) / (C_1d + 1))) / (c + eps) / Np\n        aps = aps.sum(dim=-1)\n\n        assert aps.numel() == N\n\n        if ret == '1-mAP':\n            if qw is not None:\n                aps *= qw  # query weights\n            return 1 - aps.mean()\n        elif ret == 'AP':\n            assert qw is None\n            return aps\n        else:\n            raise ValueError(\"Bad return type for APLoss(): %s\" % str(ret))\n\n    def measures(self, x, gt, loss=None):\n        if loss is None:\n            loss = self.forward(x, gt)\n        return {'loss_tap'+('s' if self.simplified else ''): float(loss)}","metadata":{},"execution_count":null,"outputs":[]}]}