You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

从零实现Contrastive Loss遇梯度爆炸,求代码错误排查

对比损失代码问题分析与修正

你的代码存在几个关键错误,直接导致了梯度爆炸和损失计算异常:

1. 余弦相似度计算完全错误

正常的余弦相似度是归一化后特征的点积,因为归一化后向量的点积等价于余弦相似度。但你当前的计算逻辑完全错误:

cosine_num = torch.matmul(z_i, z_j.T)
cosine_denom = torch.matmul(z_i_norm, z_j_norm.T)
cosine_similarity = cosine_num / cosine_denom

z_i_norm和z_j_norm本身就是L2归一化后的向量,它们的点积就是余弦相似度,不需要额外的分子分母除法。这种错误计算会产生极端数值(比如分母接近0时,相似度会趋近于无穷大),直接引发梯度爆炸。

2. 分母计算逻辑混乱

你先将denominator赋值为cosine_similarity并置空对角线,但随后又重新赋值为torch.exp(torch.sum(cosine_similarity, dim=1)),等于之前的置空操作完全无效。而且对比损失的分母应该是当前样本与所有其他样本(排除自身)的exp相似度之和,你的求和逻辑完全不符合要求。

3. 数值稳定性缺失

直接计算torch.exp再做除法,很容易因为相似度数值过大导致exp溢出,进而产生无穷大的数值,引发梯度爆炸。


修正后的对比损失代码

以下是符合SimCLR风格的对比损失实现,解决了上述所有问题:

import torch
import torch.nn as nn
import torch.nn.functional as F

class ContrastiveLoss(nn.Module):
    def __init__(self, temperature=0.5):
        super(ContrastiveLoss, self).__init__()
        self.temperature = temperature

    def forward(self, projections_1, projections_2):
        # 对特征做L2归一化
        z_i = F.normalize(projections_1, dim=1)
        z_j = F.normalize(projections_2, dim=1)
        # 合并两组特征,方便计算所有样本间的相似度
        z = torch.cat([z_i, z_j], dim=0)
        batch_size = z_i.size(0)
        
        # 计算所有样本间的余弦相似度(除以温度系数)
        cos_sim = torch.matmul(z, z.T) / self.temperature
        # 屏蔽对角线(自身与自身的相似度)
        mask = torch.eye(2 * batch_size, device=z.device, dtype=torch.bool)
        cos_sim = cos_sim.masked_fill(mask, -1e12)
        
        # 构建正样本对的掩码:z_i的正样本是z_j对应位置的样本,反之亦然
        pos_mask = torch.zeros_like(mask)
        pos_mask[:batch_size, batch_size:] = torch.eye(batch_size, device=z.device)
        pos_mask[batch_size:, :batch_size] = torch.eye(batch_size, device=z.device)
        
        # 计算损失:每个样本的损失是 -log(正样本相似度exp / 所有负样本相似度exp之和)
        numerator = torch.exp(cos_sim[pos_mask]).sum(dim=1)
        denominator = torch.exp(cos_sim).sum(dim=1)
        loss = -torch.log(numerator / denominator).sum() / (2 * batch_size)
        
        return loss

关键修正点说明

  • 正确计算归一化后的余弦相似度,并融入温度系数
  • 使用掩码区分正样本对和负样本,避免自身相似度干扰
  • 用masked_fill将自身相似度设为极小值,防止计算时包含自身
  • 最后除以2*batch_size做归一化,保证损失数值稳定
  • 优化数值计算逻辑,降低溢出风险

内容的提问来源于stack exchange,提问作者brzig

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.15 15:12:35