从零实现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
相关产品推荐
相关产品推荐

