使用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 formatnpairs_lossrequires.- 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
相关产品推荐
相关产品推荐

