基于Batch Hard Triplet Mining的孪生网络CelebA向量坍缩问题
问题:Triplet Loss在CelebA数据集上出现向量坍缩
在TensorFlow 2.10.0中复现论文时,遇到如下问题:在MNIST数据集上使用简单全连接神经网络(FCNN),搭配相同的损失函数与数据生成方式时,模型未出现坍缩;但在CelebA数据集上使用预训练ResNet50搭配两个含1024和128单元的全连接层(与论文描述一致)时,模型发生向量坍缩——所有图像的向量表示聚集在向量空间的单点,损失无法低于设定的margin值。
复现代码
class BatchHardTripletLoss(tf.keras.losses.Loss): def init(self, margin=0.5, squared=False): super().init() self.margin = margin self.squared = squared def call(self, mask, embeddings): size = tf.shape(mask)[-1]//2 mask_anchor_positive, mask_anchor_negative = mask[:, :size], mask[:, size:] pairwise_dist = self._pairwise_distances(embeddings, squared=self.squared) mask_anchor_positive = tf.cast(mask_anchor_positive, dtype=tf.float32) anchor_positive_dist = tf.multiply(mask_anchor_positive, pairwise_dist) hardest_positive_dist = tf.reduce_max(anchor_positive_dist, axis=1, keepdims=True) mask_anchor_negative = tf.cast(mask_anchor_negative, dtype=tf.float32) max_anchor_negative_dist = tf.reduce_max(pairwise_dist, axis=1, keepdims=True) anchor_negative_dist = pairwise_dist + max_anchor_negative_dist * (1.0 - mask_anchor_negative) hardest_negative_dist = tf.reduce_min(anchor_negative_dist, axis=1, keepdims=True) triplet_loss = tf.maximum((hardest_positive_dist - hardest_negative_dist) + self.margin, 0.0) triplet_loss = tf.reduce_mean(triplet_loss) return triplet_loss @staticmethod def _pairwise_distances(embeddings, squared=False): dot_product = tf.matmul(embeddings, tf.transpose(embeddings)) square_norm = tf.linalg.diag_part(dot_product) distances = tf.expand_dims(square_norm, 0) - 2.0 * dot_product + tf.expand_dims(square_norm, 1) distances = tf.maximum(distances, 0.0) if not squared: mask = tf.cast(tf.equal(distances, 0.0), dtype=tf.float32) distances = distances + mask * 1e-16 distances = tf.sqrt(distances) distances = distances * (1.0 - mask) return distances base_model = tf.keras.applications.ResNet50( include_top=False, weights="imagenet", input_shape=(224,224,3), pooling='avg' ) base_model.trainable = False # 尝试放开部分顶层权重训练,依然无效 # for layer in base_model.layers[:-8]: # layer.trainable = False inputs = tf.keras.Input(shape=(224, 224, 3)) x = base_model(inputs, training = False) x = tf.keras.layers.Dense(1024, activation='relu')(x) x = tf.keras.layers.BatchNormalization()(x) x = tf.keras.layers.ReLU()(x) output = tf.keras.layers.Dense(128)(x) embeddings = tf.keras.Model(inputs, output) class SiameseModel(tf.keras.models.Model): def __init__(self, embeddings): super().__init__() self.embeddings = embeddings self.loss_tracker = tf.keras.metrics.Mean(name="loss") def train_step(self, data): X, label = data positive_mask, negative_mask = _get_anchor_positive_triplet_mask(label), _get_anchor_negative_triplet_mask(label) mask = tf.concat((positive_mask, negative_mask),-1) with tf.GradientTape() as tape: y_pred = self(X) # 前向传播 loss = self.loss(mask, y_pred) trainable_vars = self.trainable_variables gradients = tape.gradient(loss, trainable_vars) self.optimizer.apply_gradients(zip(gradients, trainable_vars)) self.loss_tracker.update_state(loss) return {"loss": self.loss_tracker.result()} def call(self, X): return self.embeddings(X) generator = DataGenerator(eval_split[0], identity_images, img_dir, P=18, K = 4, dim=(224, 224, 3)) model.compile(optimizer = 'adam', loss = BatchHardTripletLoss(0.5))
内容的提问来源于stack exchange,提问作者Sajith 17
相关产品推荐
相关产品推荐

