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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 22:57:22