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

搭建变分自编码器时model.fit报TypeError错误求助

变分自编码器VAE训练时TypeError错误分析与解决

错误原因

报错的核心问题是自定义损失函数vae_loss直接引用了模型的中间层张量mu和log_var。Keras的Functional API要求损失函数只能接收y_true和y_pred作为入参,直接调用外部的符号张量会破坏TensorFlow的符号计算图分发逻辑,导致无法正确处理KerasTensor,触发TypeError。

解决方案

通过add_loss方法将VAE的复合损失直接整合到模型中,避免在独立损失函数中引用外部张量。具体步骤:

  • 确保输入数据维度匹配模型要求(单通道输入需扩展维度)
  • 在模型内部计算重构损失与KL损失的总和
  • 使用model.add_loss()将总损失绑定到模型,无需在编译时指定独立损失函数

修正后的完整代码

import os
import re
import numpy as np
from PIL import Image
import cv2
from sklearn.model_selection import train_test_split
from keras.layers import Input, Conv2D, Flatten, Dense, Lambda, Reshape, Conv2DTranspose
from keras.models import Model
from keras import backend as K
from keras.losses import binary_crossentropy

def load_data4():
    path_to_images = '/content/drive/MyDrive/Heatsource504'
    pattern = r'x_(\d+)y_(\d+)\.jpg'

    images = []
    heatmaps = []
    coordinates = []

    for filename in os.listdir(path_to_images):
        match = re.search(pattern, filename)
        if match:
            x_coord = int(match.group(1))
            y_coord = int(match.group(2))
            img = Image.open(os.path.join(path_to_images, filename))
            img = img.resize((200, 200))
            img_array = np.array(img)

            # 转灰度图并扩展通道维度为(200,200,1),匹配模型输入shape
            if len(img_array.shape) == 3 and img_array.shape[2] == 3:
                img_array = cv2.cvtColor(img_array, cv2.COLOR_RGB2GRAY)
            if img_array.dtype != 'uint8':
                img_array = img_array.astype('uint8')
            img_array = img_array / 255.0
            heatmaps.append(img_array[..., np.newaxis])
            
            # 原图像转RGB并归一化,保留三通道用于输出对比
            orig_img = cv2.imread(os.path.join(path_to_images, filename))
            orig_img = cv2.cvtColor(orig_img, cv2.COLOR_BGR2RGB)
            orig_img = cv2.resize(orig_img, (200, 200)) / 255.0
            images.append(orig_img)
            
            coordinates.append([x_coord, y_coord])

    images = np.array(images)
    heatmaps = np.array(heatmaps)
    coordinates = np.array(coordinates)

    X_train, X_val, y_train, y_val, coords_train, coords_val = train_test_split(heatmaps, images, coordinates, test_size=0.2, random_state=42)

    return X_train, y_train, coords_train, X_val, y_val, coords_val


input_shape = (200, 200, 1)

# Encoder
inputs_heatmap = Input(shape=input_shape)
x = Conv2D(filters=16, kernel_size=3, padding='valid', activation='relu')(inputs_heatmap)
x = Conv2D(filters=32, kernel_size=3, padding='valid', activation='relu')(x)
x = Conv2D(filters=64, kernel_size=3, padding='valid', activation='relu')(x)
x = Conv2D(filters=128, kernel_size=3, padding='valid', activation='relu')(x)
x = Conv2D(filters=256, kernel_size=3, padding='same', activation='relu')(x)
x = Flatten()(x)
x = Dense(units=128, activation='relu')(x)

# 定义潜在变量
latent_dim = 10
mu = Dense(units=latent_dim)(x)
log_var = Dense(units=latent_dim)(x)

# 重参数化技巧
def sampling(args):
    mu, log_var = args
    epsilon = K.random_normal(shape=K.shape(mu))
    return mu + K.exp(log_var / 2) * epsilon
z = Lambda(sampling)([mu, log_var])

# Decoder - 输出三通道彩色图,匹配原图像维度
x = Dense(units=128, activation='relu')(z)
x = Dense(units=8 * 8 * 128, activation='relu')(x)
x = Reshape(target_shape=(8, 8, 128))(x)
x = Conv2DTranspose(filters=128, kernel_size=3, padding='same', activation='relu')(x)
x = Conv2DTranspose(filters=64, kernel_size=4, padding='valid', activation='relu')(x)
x = Conv2DTranspose(filters=32, kernel_size=2, padding='valid', activation='relu')(x)
x = Conv2DTranspose(filters=16, kernel_size=3, padding='valid', activation='relu')(x)
outputs = Conv2DTranspose(filters=3, kernel_size=2, padding='valid', activation='sigmoid')(x)

# 多输入模型:输入heatmap和原图像(用于计算重构损失)
inputs_target = Input(shape=(200,200,3))
model = Model(inputs=[inputs_heatmap, inputs_target], outputs=outputs)

# 计算VAE总损失
# 重构损失:基于原图像和模型输出计算
reconstruction_loss = binary_crossentropy(K.flatten(inputs_target), K.flatten(outputs))
reconstruction_loss *= 200 * 200 * 3  # 按像素数加权
# KL损失:基于潜在变量的分布计算
kl_loss = 1 + log_var - K.square(mu) - K.exp(log_var)
kl_loss = K.sum(kl_loss, axis=-1)
kl_loss *= -0.5
total_loss = K.mean(reconstruction_loss + kl_loss)

# 将损失绑定到模型
model.add_loss(total_loss)

# 编译模型,无需指定loss参数
model.compile(optimizer='adam')

# 加载数据并训练
X_train, y_train, coords_train, X_val, y_val, coords_val = load_data4()
model.fit(
    [X_train, y_train],
    epochs=10,
    batch_size=16,
    validation_data=([X_val, y_val], None)
)

关键调整说明

  1. 数据维度修正:给灰度heatmap添加通道维度,确保与模型输入(200,200,1)匹配;原图像统一转RGB格式,匹配模型输出的三通道维度
  2. 多输入模型构建:新增原图像输入张量,用于计算重构损失,避免损失函数引用外部变量
  3. 损失绑定方式:用model.add_loss()将复合损失直接整合到模型,符合Keras符号计算规范
  4. 输出通道调整:解码器最后一层设为3个滤波器,对应彩色图像输出

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 05:44:55