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

使用Keras实现监督对比学习时因独热编码标签引发InvalidArgumentError问题求助

Fixing Supervised Contrastive Loss with One-Hot Encoded Labels

The error you're encountering stems from a mismatch in label format: tfa.losses.npairs_loss expects integer (sparse) labels (shape (batch_size,)) but you’re passing one-hot encoded labels (shape (batch_size, num_classes)). The tf.squeeze(labels) call doesn’t resolve this because your labels don’t have a singleton dimension to squeeze out.

The Solution: Convert One-Hot Labels to Integer Labels

Modify your SupervisedContrastiveLoss class to convert one-hot labels to integer indices using tf.argmax before passing them to npairs_loss. Here’s the corrected code:

class SupervisedContrastiveLoss(keras.losses.Loss):
    def __init__(self, temperature=1, name=None):
        super(SupervisedContrastiveLoss, self).__init__(name=name)
        self.temperature = temperature

    def __call__(self, labels, feature_vectors, sample_weight=None):
        # Normalize feature vectors
        feature_vectors_normalized = tf.math.l2_normalize(feature_vectors, axis=1)
        # Compute logits
        logits = tf.divide(
            tf.matmul(
                feature_vectors_normalized, tf.transpose(feature_vectors_normalized)
            ),
            self.temperature,
        )
        # Convert one-hot labels to integer labels (shape from (batch_size, 2) to (batch_size,))
        integer_labels = tf.argmax(labels, axis=1)
        return tfa.losses.npairs_loss(integer_labels, logits)

Why This Works

  • tf.argmax(labels, axis=1) transforms your one-hot labels (e.g., [[1,0], [0,1], [1,0]]) into integer indices ([0, 1, 0]), which matches the format npairs_loss requires.
  • This resolves the shape mismatch between logits ([batch_size, batch_size]) and labels, allowing the loss calculation to run without errors.

Optional Debug Check

To confirm the label shape conversion is working as expected, add a debug line inside the __call__ method:

tf.print("Integer labels shape:", tf.shape(integer_labels))  # Should output (batch_size,)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.29 21:24:07