Keras多输入UNet自定义损失函数问题:训练后输出与输入一致
问题描述
基于Keras构建带权重图输入的UNet模型(参考原始UNet论文)用于图像合成任务,损失函数结合感知损失与像素损失,需用到输入图像、重建图像和权重图三个输入。训练时通过loss=None完成模型编译,但预测结果与输入图像完全一致,网络未产生任何有效修改。
用户提供的网络与损失函数代码如下:
def synthesis_unet_weights(pretrained_weights=None, input_shape=(SIZE_s, SIZE_s, 3), num_classes=1, is_training=True): ip = Input(shape=input_shape) weight_ip = Input(shape=input_shape[:2] + (num_classes,)) # UNET encoder with the first Conv2D layer taking input ip #--------------------------------------------------------------------------------------------------------------------------- center = Conv2D(1024, (3,3),padding='same', activation='relu', kernel_initializer=initializer)(pool4) center = Conv2D(1024, (3,3),padding='same', activation='relu', kernel_initializer=initializer)(center) #--------------------------------------------------------------------------------------------------------------------------- # UNET decoder with the last layer up1 classify = Conv2D(num_classes, (1,1), activation='sigmoid')(up1) if is_training: model=Model(inputs=[ip, weight_ip], outputs=[classify]) model.add_loss(perceptual_loss_weight(ip,classify,weight_ip)) return model else: model = Model(inputs=[ip], outputs=[classify]) weight_ip=ip model.add_loss(perceptual_loss_weight(ip,classify,weight_ip)) opt2 = tf.keras.optimizers.Adam(learning_rate=1e-3,clipnorm=1.0) model.compile(optimizer=opt2) return model return model def perceptual_loss_weight(input_image , reconstruct_image, weights): input_image = clip_0_1(input_image) reconstruct_image = tf.concat((reconstruct_image,reconstruct_image,reconstruct_image),axis=-1) reconstruct_image = clip_0_1(reconstruct_image) weights = tf.concat((weights,weights,weights),axis=-1) weights = clip_0_1(weights) h1_list = LossModel(input_image) h2_list = LossModel(reconstruct_image) rc_loss = 0.0 for h1, h2, weight in zip(h1_list, h2_list, selected_layer_weights): h1 = K.batch_flatten(h1) h2 = K.batch_flatten(h2) rc_loss = rc_loss + weight * K.sum(K.square(h1 - h2), axis=-1) pixel_loss = K.sum(K.square(K.batch_flatten(weights)*K.batch_flatten(input_image) - K.batch_flatten(weights)*K.batch_flatten(reconstruct_image)),axis=1) return rc_loss+pixel_loss
排查与修复建议
预测阶段模型逻辑错误:
is_training=False分支中,无需调用model.add_loss,预测阶段仅需前向传播输出结果,绑定损失会干扰计算图。同时错误地将weight_ip=ip完全没必要,应修改为:else: model = Model(inputs=[ip], outputs=[classify]) # 预测阶段无需计算损失,直接返回模型 return model训练阶段未编译模型:训练分支返回的模型未配置优化器,即便用
add_loss定义了损失,也无法更新权重。修改训练分支:if is_training: model=Model(inputs=[ip, weight_ip], outputs=[classify]) model.add_loss(perceptual_loss_weight(ip,classify,weight_ip)) opt = tf.keras.optimizers.Adam(learning_rate=1e-4) # 1e-3学习率过高,建议下调 model.compile(optimizer=opt) return model像素损失计算逻辑有误:当前展平操作破坏了空间维度的权重对应关系,应保留空间维度计算加权误差:
# 替换原pixel_loss计算 pixel_loss = K.sum(K.square(weights * (input_image - reconstruct_image)), axis=[1,2,3])原写法将所有通道展平后相乘,会导致权重与像素的对应关系错乱,无法正确发挥权重图的区域约束作用。
UNet主体结构完整性检查:代码中编码器、解码器为占位符,需确认实际实现的结构是否正确:
- 编码器需包含正确的下采样(MaxPooling2D)和卷积层堆叠
- 解码器需包含上采样(UpSampling2D或Conv2DTranspose)与编码器特征图的拼接(Concatenate)操作
- 若编码器和解码器未正确连接,比如直接将输入传递到输出,必然导致输出与输入一致
权重图有效性验证:训练时需确认传入的
weight_ip是否符合预期:- 权重图是否存在非1/0的有效权重分布,若全为1则等价于普通像素损失,若全为0则像素损失失效
- 检查权重图的维度是否与输入图像匹配,拼接操作是否正确
LossModel与感知损失验证:
- 确认
LossModel是预训练的特征提取模型(如VGG)且已冻结权重,若模型未正确提取特征或返回输入本身,感知损失会驱动模型拟合输入 - 检查
selected_layer_weights是否为正的有效权重值,若全为0则感知损失不起作用,仅剩像素损失
- 确认
内容的提问来源于stack exchange,提问作者ML_student
相关产品推荐
相关产品推荐

