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

实现InfoNCE损失时出现张量形状不匹配问题求助

问题:InfoNCE损失调整batch_size时出现维度不匹配错误

背景

我基于InfoNCE相关论文在自定义图像数据集上实现自监督训练,参考SimCLR仓库代码编写了InfoNCE损失函数:

def info_nce_loss(self, features):

        labels = torch.cat([torch.arange(self.args.batch_size) for i in range(self.args.n_views)], dim=0)
        labels = (labels.unsqueeze(0) == labels.unsqueeze(1)).float()
        labels = labels.to(self.args.device)

        features = F.normalize(features, dim=1)

        similarity_matrix = torch.matmul(features, features.T)
        # assert similarity_matrix.shape == (
        #     self.args.n_views * self.args.batch_size, self.args.n_views * self.args.batch_size)
        # assert similarity_matrix.shape == labels.shape

        # discard the main diagonal from both: labels and similarities matrix
        mask = torch.eye(labels.shape[0], dtype=torch.bool).to(self.args.device)
        labels = labels[~mask].view(labels.shape[0], -1)
        similarity_matrix = similarity_matrix[~mask].view(similarity_matrix.shape[0], -1)
        # assert similarity_matrix.shape == labels.shape

        # select and combine multiple positives
        positives = similarity_matrix[labels.bool()].view(labels.shape[0], -1)

        # select only the negatives the negatives
        negatives = similarity_matrix[~labels.bool()].view(similarity_matrix.shape[0], -1)

        logits = torch.cat([positives, negatives], dim=1)
        labels = torch.zeros(logits.shape[0], dtype=torch.long).to(self.args.device)

        logits = logits / self.args.temperature
        return logits, labels

当batch_size设为32时训练正常,但改为256等数值时,在labels = labels[~mask].view(labels.shape[0], -1)行抛出错误:

The shape of the mask [512, 512] at index 0 does not match the shape of the indexed tensor [2, 2] at index 0. 

调整图像尺寸无法解决问题,请问问题原因是什么?


问题原因分析

这个错误的核心是实际训练使用的batch_size与self.args.batch_size参数不一致,导致张量维度完全冲突:

  1. 当batch_size改为256时,输入到损失函数的features张量维度应为[256×n_views, feature_dim](假设n_views=2,即[512, D]),对应的similarity_matrix是512×512的矩阵。
  2. 但self.args.batch_size并未同步更新(仍保留初始的极小值,比如2),导致生成的labels张量维度为[2×n_views, 2×n_views](即2×2)。
  3. 代码中mask是基于similarity_matrix的维度生成的512×512张量,用它去索引2×2的labels时,必然出现维度不匹配的错误。
  4. 你注释掉了原代码中assert similarity_matrix.shape == labels.shape的断言,错过了提前发现维度不匹配的机会,导致问题延后到索引操作时才暴露。

此外,也可能是多GPU训练场景下的参数配置错误:如果self.args.batch_size设为全局batch_size,但单GPU上的实际batch_size是全局值除以GPU数量,会导致labels生成的维度远小于features的维度,引发同样的错误。


解决办法

  1. 同步参数与实际配置:确保DataLoader设置的batch_size和self.args.batch_size完全一致,修改batch_size时同步更新self.args中的对应参数。
  2. 恢复断言检查:取消注释原代码中的维度断言,这样可以在维度不匹配时立刻报错,快速定位问题。
  3. 适配多GPU训练:如果使用分布式训练,self.args.batch_size应设为单GPU的local batch_size,而非全局batch_size;或使用分布式张量生成匹配维度的labels。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 07:59:55