{
  "id": 505728,
  "title": "[HELP] Supervised contrastive loss exploding",
  "url": "/competitions/birdclef-2024/discussion/505728",
  "author_name": "Arindam Roy",
  "post_date": "2024-05-18T19:50:14.839000",
  "votes": 1,
  "comment_count": 0,
  "views": 0,
  "content": "<p>Hey! </p>\n<p>So I think supervised contrastive loss would be a great idea for this comp. Hence I tried implementing this <a href=\"https://arxiv.org/pdf/2004.11362.pdf\" target=\"_blank\">paper</a>. However, I have no clue why, but instead of reducing, the contrastive loss shoots up with every batch till I get a Nan. I have experimented with different LR, batch sizes, loss function implementation but same result. I wanted to ask if anyone has made similar tries and has encountered similar problems? </p>\n<p><a href=\"https://github.com/HobbitLong/SupContrast/tree/master\" target=\"_blank\">Loss</a> :<br>\n`<br>\nclass SupConLoss(nn.Module):<br>\n    \"\"\"Supervised Contrastive Learning: <a href=\"https://arxiv.org/pdf/2004.11362.pdf\" target=\"_blank\">https://arxiv.org/pdf/2004.11362.pdf</a>.<br>\n    It also supports the unsupervised contrastive loss in SimCLR\"\"\"<br>\n    def <strong>init</strong>(self, temperature=0.07, contrast_mode='all',<br>\n                 base_temperature=0.07):<br>\n        super(SupConLoss, self).<strong>init</strong>()<br>\n        self.temperature = temperature<br>\n        self.contrast_mode = contrast_mode<br>\n        self.base_temperature = base_temperature</p>\n<pre><code> ():\n    \n     (labels.shape) == :\n        labels = torch.argmax(labels, dim=)\n\n\n    device = (torch.device()\n               features.is_cuda\n               torch.device())\n\n     (features.shape) &lt; :\n         ValueError(\n                         )\n     (features.shape) &gt; :\n        features = features.view(features.shape[], features.shape[], -)\n\n    batch_size = features.shape[]\n     labels     mask   :\n         ValueError()\n     labels    mask  :\n        mask = torch.eye(batch_size, dtype=torch.float32).to(device)\n     labels   :\n        labels = labels.contiguous().view(-, )\n         labels.shape[] != batch_size:\n             ValueError()\n        mask = torch.eq(labels, labels.T).().to(device)\n    :\n        mask = mask.().to(device)\n\n    contrast_count = features.shape[]\n    contrast_feature = torch.cat(torch.unbind(features, dim=), dim=)\n     self.contrast_mode == :\n        anchor_feature = features[:, ]\n        anchor_count = \n     self.contrast_mode == :\n        anchor_feature = contrast_feature\n        anchor_count = contrast_count\n    :\n         ValueError(.(self.contrast_mode))\n\n    \n    anchor_dot_contrast = torch.div(\n        torch.matmul(anchor_feature, contrast_feature.T),\n        self.temperature)\n    \n    logits_max, _ = torch.(anchor_dot_contrast, dim=, keepdim=)\n    logits = anchor_dot_contrast - logits_max.detach()\n\n    \n    mask = mask.repeat(anchor_count, contrast_count)\n    \n    logits_mask = torch.scatter(\n        torch.ones_like(mask),\n        ,\n        torch.arange(batch_size * anchor_count).view(-, ).to(device),\n        \n    )\n    mask = mask * logits_mask\n\n    \n    exp_logits = torch.exp(logits) * logits_mask\n    log_prob = logits - torch.log(exp_logits.(, keepdim=))\n\n    \n    \n    \n    \n    \n    \n    \n    mask_pos_pairs = mask.()\n    mask_pos_pairs = torch.where(mask_pos_pairs &lt; , , mask_pos_pairs)\n    mean_log_prob_pos = (mask * log_prob).() / mask_pos_pairs\n\n    \n    loss = - (self.temperature / self.base_temperature) * mean_log_prob_pos\n    loss = loss.view(anchor_count, batch_size).mean()\n\n     loss\n</code></pre>\n<p>`</p>\n<p>Dataset:<br>\n`</p>\n<p>class BirdDataset(torch.utils.data.Dataset):</p>\n<pre><code>def __init__(self, df, sr = Config.SR, \n             duration = Config.DURATION, augmentations = None, \n             train = , transform = transform):\n\n    self.df = df\n    self.sr = sr \n    self.train = train\n    self.duration = duration\n    self.augmentations = augmentations\n    self.transform = transform\n\n     train:\n        self.img_dir = Config.train_images\n    else:\n        self.img_dir = Config.valid_images\n\ndef __len__(self):\n    return len(self.df)\n\n@staticmethod\ndef normalize(image):\n    image = image / \n    #image = torch.stack([image, image, image])\n    return image\n\n\n\ndef __getitem__(self, idx):\n\n    row = self.df.iloc[idx]\n    impath = self.img_dir + f\n\n    image = np.load(str(impath))[:Config.MAX_READ_SAMPLES]\n\n    ########## RANDOM SAMPLING ################\n     self.train:\n        image = image[np.random.choice(len(image))]\n    else:\n        image = image[]\n\n    #####################################################################\n\n    image = torch.tensor(image).float()\n\n     self.augmentations:\n        image = self.augmentations(image.unsqueeze()).squeeze()\n     self.transform:\n        image_pil = transforms.ToPILImage()(image)\n        image= self.transform(image_pil)  # Apply transformations\n        image = transforms.ToTensor()(image)\n\n    image.size()\n\n\n     self.augmentations:\n        image2 = self.augmentations(image2.unsqueeze()).squeeze()\n     self.transform:\n        image2 = self.transform(image_pil)  # Apply transformations\n        image2 = transforms.ToTensor()(image2)\n\n    image = self.normalize(image)\n    image2 = self.normalize(image)\n    label = torch.tensor(row[:]).float()\n    label = torch.argmax(label, dim=)\n\n\n    return (image, image2) , label\n</code></pre>\n<p>`</p>",
  "messages": [
    {
      "id": 2822842,
      "postDate": "2024-05-18T19:50:14.840Z",
      "content": "<p>Hey! </p>\n<p>So I think supervised contrastive loss would be a great idea for this comp. Hence I tried implementing this <a href=\"https://arxiv.org/pdf/2004.11362.pdf\" target=\"_blank\">paper</a>. However, I have no clue why, but instead of reducing, the contrastive loss shoots up with every batch till I get a Nan. I have experimented with different LR, batch sizes, loss function implementation but same result. I wanted to ask if anyone has made similar tries and has encountered similar problems? </p>\n<p><a href=\"https://github.com/HobbitLong/SupContrast/tree/master\" target=\"_blank\">Loss</a> :<br>\n`<br>\nclass SupConLoss(nn.Module):<br>\n    \"\"\"Supervised Contrastive Learning: <a href=\"https://arxiv.org/pdf/2004.11362.pdf\" target=\"_blank\">https://arxiv.org/pdf/2004.11362.pdf</a>.<br>\n    It also supports the unsupervised contrastive loss in SimCLR\"\"\"<br>\n    def <strong>init</strong>(self, temperature=0.07, contrast_mode='all',<br>\n                 base_temperature=0.07):<br>\n        super(SupConLoss, self).<strong>init</strong>()<br>\n        self.temperature = temperature<br>\n        self.contrast_mode = contrast_mode<br>\n        self.base_temperature = base_temperature</p>\n<pre><code> ():\n    \n     (labels.shape) == :\n        labels = torch.argmax(labels, dim=)\n\n\n    device = (torch.device()\n               features.is_cuda\n               torch.device())\n\n     (features.shape) &lt; :\n         ValueError(\n                         )\n     (features.shape) &gt; :\n        features = features.view(features.shape[], features.shape[], -)\n\n    batch_size = features.shape[]\n     labels     mask   :\n         ValueError()\n     labels    mask  :\n        mask = torch.eye(batch_size, dtype=torch.float32).to(device)\n     labels   :\n        labels = labels.contiguous().view(-, )\n         labels.shape[] != batch_size:\n             ValueError()\n        mask = torch.eq(labels, labels.T).().to(device)\n    :\n        mask = mask.().to(device)\n\n    contrast_count = features.shape[]\n    contrast_feature = torch.cat(torch.unbind(features, dim=), dim=)\n     self.contrast_mode == :\n        anchor_feature = features[:, ]\n        anchor_count = \n     self.contrast_mode == :\n        anchor_feature = contrast_feature\n        anchor_count = contrast_count\n    :\n         ValueError(.(self.contrast_mode))\n\n    \n    anchor_dot_contrast = torch.div(\n        torch.matmul(anchor_feature, contrast_feature.T),\n        self.temperature)\n    \n    logits_max, _ = torch.(anchor_dot_contrast, dim=, keepdim=)\n    logits = anchor_dot_contrast - logits_max.detach()\n\n    \n    mask = mask.repeat(anchor_count, contrast_count)\n    \n    logits_mask = torch.scatter(\n        torch.ones_like(mask),\n        ,\n        torch.arange(batch_size * anchor_count).view(-, ).to(device),\n        \n    )\n    mask = mask * logits_mask\n\n    \n    exp_logits = torch.exp(logits) * logits_mask\n    log_prob = logits - torch.log(exp_logits.(, keepdim=))\n\n    \n    \n    \n    \n    \n    \n    \n    mask_pos_pairs = mask.()\n    mask_pos_pairs = torch.where(mask_pos_pairs &lt; , , mask_pos_pairs)\n    mean_log_prob_pos = (mask * log_prob).() / mask_pos_pairs\n\n    \n    loss = - (self.temperature / self.base_temperature) * mean_log_prob_pos\n    loss = loss.view(anchor_count, batch_size).mean()\n\n     loss\n</code></pre>\n<p>`</p>\n<p>Dataset:<br>\n`</p>\n<p>class BirdDataset(torch.utils.data.Dataset):</p>\n<pre><code>def __init__(self, df, sr = Config.SR, \n             duration = Config.DURATION, augmentations = None, \n             train = , transform = transform):\n\n    self.df = df\n    self.sr = sr \n    self.train = train\n    self.duration = duration\n    self.augmentations = augmentations\n    self.transform = transform\n\n     train:\n        self.img_dir = Config.train_images\n    else:\n        self.img_dir = Config.valid_images\n\ndef __len__(self):\n    return len(self.df)\n\n@staticmethod\ndef normalize(image):\n    image = image / \n    #image = torch.stack([image, image, image])\n    return image\n\n\n\ndef __getitem__(self, idx):\n\n    row = self.df.iloc[idx]\n    impath = self.img_dir + f\n\n    image = np.load(str(impath))[:Config.MAX_READ_SAMPLES]\n\n    ########## RANDOM SAMPLING ################\n     self.train:\n        image = image[np.random.choice(len(image))]\n    else:\n        image = image[]\n\n    #####################################################################\n\n    image = torch.tensor(image).float()\n\n     self.augmentations:\n        image = self.augmentations(image.unsqueeze()).squeeze()\n     self.transform:\n        image_pil = transforms.ToPILImage()(image)\n        image= self.transform(image_pil)  # Apply transformations\n        image = transforms.ToTensor()(image)\n\n    image.size()\n\n\n     self.augmentations:\n        image2 = self.augmentations(image2.unsqueeze()).squeeze()\n     self.transform:\n        image2 = self.transform(image_pil)  # Apply transformations\n        image2 = transforms.ToTensor()(image2)\n\n    image = self.normalize(image)\n    image2 = self.normalize(image)\n    label = torch.tensor(row[:]).float()\n    label = torch.argmax(label, dim=)\n\n\n    return (image, image2) , label\n</code></pre>\n<p>`</p>",
      "rawMarkdown": "Hey! \n\nSo I think supervised contrastive loss would be a great idea for this comp. Hence I tried implementing this [paper](https://arxiv.org/pdf/2004.11362.pdf). However, I have no clue why, but instead of reducing, the contrastive loss shoots up with every batch till I get a Nan. I have experimented with different LR, batch sizes, loss function implementation but same result. I wanted to ask if anyone has made similar tries and has encountered similar problems? \n\n[Loss](https://github.com/HobbitLong/SupContrast/tree/master) :\n`\nclass SupConLoss(nn.Module):\n    \"\"\"Supervised Contrastive Learning: https://arxiv.org/pdf/2004.11362.pdf.\n    It also supports the unsupervised contrastive loss in SimCLR\"\"\"\n    def __init__(self, temperature=0.07, contrast_mode='all',\n                 base_temperature=0.07):\n        super(SupConLoss, self).__init__()\n        self.temperature = temperature\n        self.contrast_mode = contrast_mode\n        self.base_temperature = base_temperature\n\n    def forward(self, features, labels=None, mask=None):\n        \"\"\"Compute loss for model. If both `labels` and `mask` are None,\n        it degenerates to SimCLR unsupervised loss:\n        https://arxiv.org/pdf/2002.05709.pdf\n\n        Args:\n            features: hidden vector of shape [bsz, n_views, ...].\n            labels: ground truth of shape [bsz].\n            mask: contrastive mask of shape [bsz, bsz], mask_{i,j}=1 if sample j\n                has the same class as sample i. Can be asymmetric.\n        Returns:\n            A loss scalar.\n        \"\"\"\n        if len(labels.shape) == 2:\n            labels = torch.argmax(labels, dim=1)\n        \n        \n        device = (torch.device('cuda')\n                  if features.is_cuda\n                  else torch.device('cpu'))\n\n        if len(features.shape) < 3:\n            raise ValueError('`features` needs to be [bsz, n_views, ...],'\n                             'at least 3 dimensions are required')\n        if len(features.shape) > 3:\n            features = features.view(features.shape[0], features.shape[1], -1)\n\n        batch_size = features.shape[0]\n        if labels is not None and mask is not None:\n            raise ValueError('Cannot define both `labels` and `mask`')\n        elif labels is None and mask is None:\n            mask = torch.eye(batch_size, dtype=torch.float32).to(device)\n        elif labels is not None:\n            labels = labels.contiguous().view(-1, 1)\n            if labels.shape[0] != batch_size:\n                raise ValueError('Num of labels does not match num of features')\n            mask = torch.eq(labels, labels.T).float().to(device)\n        else:\n            mask = mask.float().to(device)\n\n        contrast_count = features.shape[1]\n        contrast_feature = torch.cat(torch.unbind(features, dim=1), dim=0)\n        if self.contrast_mode == 'one':\n            anchor_feature = features[:, 0]\n            anchor_count = 1\n        elif self.contrast_mode == 'all':\n            anchor_feature = contrast_feature\n            anchor_count = contrast_count\n        else:\n            raise ValueError('Unknown mode: {}'.format(self.contrast_mode))\n\n        # compute logits\n        anchor_dot_contrast = torch.div(\n            torch.matmul(anchor_feature, contrast_feature.T),\n            self.temperature)\n        # for numerical stability\n        logits_max, _ = torch.max(anchor_dot_contrast, dim=1, keepdim=True)\n        logits = anchor_dot_contrast - logits_max.detach()\n\n        # tile mask\n        mask = mask.repeat(anchor_count, contrast_count)\n        # mask-out self-contrast cases\n        logits_mask = torch.scatter(\n            torch.ones_like(mask),\n            1,\n            torch.arange(batch_size * anchor_count).view(-1, 1).to(device),\n            0\n        )\n        mask = mask * logits_mask\n\n        # compute log_prob\n        exp_logits = torch.exp(logits) * logits_mask\n        log_prob = logits - torch.log(exp_logits.sum(1, keepdim=True))\n\n        # compute mean of log-likelihood over positive\n        # modified to handle edge cases when there is no positive pair\n        # for an anchor point. \n        # Edge case e.g.:- \n        # features of shape: [4,1,...]\n        # labels:            [0,1,1,2]\n        # loss before mean:  [nan, ..., ..., nan] \n        mask_pos_pairs = mask.sum(1)\n        mask_pos_pairs = torch.where(mask_pos_pairs < 1e-6, 1, mask_pos_pairs)\n        mean_log_prob_pos = (mask * log_prob).sum(1) / mask_pos_pairs\n\n        # loss\n        loss = - (self.temperature / self.base_temperature) * mean_log_prob_pos\n        loss = loss.view(anchor_count, batch_size).mean()\n\n        return loss\n`\n\nDataset:\n`\n\nclass BirdDataset(torch.utils.data.Dataset):\n\n    def __init__(self, df, sr = Config.SR, \n                 duration = Config.DURATION, augmentations = None, \n                 train = True, transform = transform):\n\n        self.df = df\n        self.sr = sr \n        self.train = train\n        self.duration = duration\n        self.augmentations = augmentations\n        self.transform = transform\n\n        if train:\n            self.img_dir = Config.train_images\n        else:\n            self.img_dir = Config.valid_images\n\n    def __len__(self):\n        return len(self.df)\n\n    @staticmethod\n    def normalize(image):\n        image = image / 255.0\n        #image = torch.stack([image, image, image])\n        return image\n    \n    \n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n        impath = self.img_dir + f\"{row.filename}.npy\"\n\n        image = np.load(str(impath))[:Config.MAX_READ_SAMPLES]\n        \n        ########## RANDOM SAMPLING ################\n        if self.train:\n            image = image[np.random.choice(len(image))]\n        else:\n            image = image[0]\n            \n        #####################################################################\n        \n        image = torch.tensor(image).float()\n\n        if self.augmentations:\n            image = self.augmentations(image.unsqueeze(0)).squeeze()\n        if self.transform:\n            image_pil = transforms.ToPILImage()(image)\n            image= self.transform(image_pil)  # Apply transformations\n            image = transforms.ToTensor()(image)\n            \n        image.size()\n        \n        \n        if self.augmentations:\n            image2 = self.augmentations(image2.unsqueeze(0)).squeeze()\n        if self.transform:\n            image2 = self.transform(image_pil)  # Apply transformations\n            image2 = transforms.ToTensor()(image2)\n\n        image = self.normalize(image)\n        image2 = self.normalize(image)\n        label = torch.tensor(row[14:]).float()\n        label = torch.argmax(label, dim=1)\n\n\n        return (image, image2) , label\n\n`\n",
      "votes": 1
    }
  ],
  "comments": [],
  "raw_markdown_by_id": {
    "2822842": "Hey! \n\nSo I think supervised contrastive loss would be a great idea for this comp. Hence I tried implementing this [paper](https://arxiv.org/pdf/2004.11362.pdf). However, I have no clue why, but instead of reducing, the contrastive loss shoots up with every batch till I get a Nan. I have experimented with different LR, batch sizes, loss function implementation but same result. I wanted to ask if anyone has made similar tries and has encountered similar problems? \n\n[Loss](https://github.com/HobbitLong/SupContrast/tree/master) :\n`\nclass SupConLoss(nn.Module):\n    \"\"\"Supervised Contrastive Learning: https://arxiv.org/pdf/2004.11362.pdf.\n    It also supports the unsupervised contrastive loss in SimCLR\"\"\"\n    def __init__(self, temperature=0.07, contrast_mode='all',\n                 base_temperature=0.07):\n        super(SupConLoss, self).__init__()\n        self.temperature = temperature\n        self.contrast_mode = contrast_mode\n        self.base_temperature = base_temperature\n\n    def forward(self, features, labels=None, mask=None):\n        \"\"\"Compute loss for model. If both `labels` and `mask` are None,\n        it degenerates to SimCLR unsupervised loss:\n        https://arxiv.org/pdf/2002.05709.pdf\n\n        Args:\n            features: hidden vector of shape [bsz, n_views, ...].\n            labels: ground truth of shape [bsz].\n            mask: contrastive mask of shape [bsz, bsz], mask_{i,j}=1 if sample j\n                has the same class as sample i. Can be asymmetric.\n        Returns:\n            A loss scalar.\n        \"\"\"\n        if len(labels.shape) == 2:\n            labels = torch.argmax(labels, dim=1)\n        \n        \n        device = (torch.device('cuda')\n                  if features.is_cuda\n                  else torch.device('cpu'))\n\n        if len(features.shape) < 3:\n            raise ValueError('`features` needs to be [bsz, n_views, ...],'\n                             'at least 3 dimensions are required')\n        if len(features.shape) > 3:\n            features = features.view(features.shape[0], features.shape[1], -1)\n\n        batch_size = features.shape[0]\n        if labels is not None and mask is not None:\n            raise ValueError('Cannot define both `labels` and `mask`')\n        elif labels is None and mask is None:\n            mask = torch.eye(batch_size, dtype=torch.float32).to(device)\n        elif labels is not None:\n            labels = labels.contiguous().view(-1, 1)\n            if labels.shape[0] != batch_size:\n                raise ValueError('Num of labels does not match num of features')\n            mask = torch.eq(labels, labels.T).float().to(device)\n        else:\n            mask = mask.float().to(device)\n\n        contrast_count = features.shape[1]\n        contrast_feature = torch.cat(torch.unbind(features, dim=1), dim=0)\n        if self.contrast_mode == 'one':\n            anchor_feature = features[:, 0]\n            anchor_count = 1\n        elif self.contrast_mode == 'all':\n            anchor_feature = contrast_feature\n            anchor_count = contrast_count\n        else:\n            raise ValueError('Unknown mode: {}'.format(self.contrast_mode))\n\n        # compute logits\n        anchor_dot_contrast = torch.div(\n            torch.matmul(anchor_feature, contrast_feature.T),\n            self.temperature)\n        # for numerical stability\n        logits_max, _ = torch.max(anchor_dot_contrast, dim=1, keepdim=True)\n        logits = anchor_dot_contrast - logits_max.detach()\n\n        # tile mask\n        mask = mask.repeat(anchor_count, contrast_count)\n        # mask-out self-contrast cases\n        logits_mask = torch.scatter(\n            torch.ones_like(mask),\n            1,\n            torch.arange(batch_size * anchor_count).view(-1, 1).to(device),\n            0\n        )\n        mask = mask * logits_mask\n\n        # compute log_prob\n        exp_logits = torch.exp(logits) * logits_mask\n        log_prob = logits - torch.log(exp_logits.sum(1, keepdim=True))\n\n        # compute mean of log-likelihood over positive\n        # modified to handle edge cases when there is no positive pair\n        # for an anchor point. \n        # Edge case e.g.:- \n        # features of shape: [4,1,...]\n        # labels:            [0,1,1,2]\n        # loss before mean:  [nan, ..., ..., nan] \n        mask_pos_pairs = mask.sum(1)\n        mask_pos_pairs = torch.where(mask_pos_pairs < 1e-6, 1, mask_pos_pairs)\n        mean_log_prob_pos = (mask * log_prob).sum(1) / mask_pos_pairs\n\n        # loss\n        loss = - (self.temperature / self.base_temperature) * mean_log_prob_pos\n        loss = loss.view(anchor_count, batch_size).mean()\n\n        return loss\n`\n\nDataset:\n`\n\nclass BirdDataset(torch.utils.data.Dataset):\n\n    def __init__(self, df, sr = Config.SR, \n                 duration = Config.DURATION, augmentations = None, \n                 train = True, transform = transform):\n\n        self.df = df\n        self.sr = sr \n        self.train = train\n        self.duration = duration\n        self.augmentations = augmentations\n        self.transform = transform\n\n        if train:\n            self.img_dir = Config.train_images\n        else:\n            self.img_dir = Config.valid_images\n\n    def __len__(self):\n        return len(self.df)\n\n    @staticmethod\n    def normalize(image):\n        image = image / 255.0\n        #image = torch.stack([image, image, image])\n        return image\n    \n    \n\n    def __getitem__(self, idx):\n\n        row = self.df.iloc[idx]\n        impath = self.img_dir + f\"{row.filename}.npy\"\n\n        image = np.load(str(impath))[:Config.MAX_READ_SAMPLES]\n        \n        ########## RANDOM SAMPLING ################\n        if self.train:\n            image = image[np.random.choice(len(image))]\n        else:\n            image = image[0]\n            \n        #####################################################################\n        \n        image = torch.tensor(image).float()\n\n        if self.augmentations:\n            image = self.augmentations(image.unsqueeze(0)).squeeze()\n        if self.transform:\n            image_pil = transforms.ToPILImage()(image)\n            image= self.transform(image_pil)  # Apply transformations\n            image = transforms.ToTensor()(image)\n            \n        image.size()\n        \n        \n        if self.augmentations:\n            image2 = self.augmentations(image2.unsqueeze(0)).squeeze()\n        if self.transform:\n            image2 = self.transform(image_pil)  # Apply transformations\n            image2 = transforms.ToTensor()(image2)\n\n        image = self.normalize(image)\n        image2 = self.normalize(image)\n        label = torch.tensor(row[14:]).float()\n        label = torch.argmax(label, dim=1)\n\n\n        return (image, image2) , label\n\n`\n"
  }
}