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

如何解决KerasTensor传入TensorFlow API引发的TypeError?WGAN-GP代码适配TF Keras调试求助

解决WGAN-GP适配TensorFlow Keras时的KerasTensor报错问题

这个报错的核心原因是旧Keras代码中的梯度计算逻辑不符合TensorFlow 2.x Keras Functional API的规范:你直接在loss函数中对KerasTensor(符号张量)调用了K.gradients,而TF2的Keras不允许在自定义loss中直接使用这类不支持符号张量调度的TF API。下面是具体的解决方案和完整适配后的代码:

关键问题分析

报错提示明确指出:不能将KerasTensor传给tf.gradients这类不支持自定义调度的API。原来的代码通过partial把interpolated_img(KerasTensor)传入loss函数,在loss里计算梯度,这种方式在TF2的Functional模式下是不兼容的,因为符号张量的梯度计算需要被封装在tf.GradientTape或者自定义层的call方法中。

解决方案步骤

  1. 修复自定义层的拼写错误:RandomWeightedAverage层的comput_output_shape方法名拼写错误,改为compute_output_shape,确保Keras能正确推断输出形状。
  2. 改用自定义训练循环:放弃原来的model.compile方式,用tf.GradientTape手动计算并应用梯度,这是TF2中处理复杂训练逻辑(如梯度惩罚)的标准方式。
  3. 重构梯度惩罚计算:把梯度惩罚的逻辑封装到单独的方法中,用GradientTape记录插值图像的梯度,避免直接在loss函数中处理符号张量。
  4. 适配TF2 API细节:比如优化器的lr参数改为learning_rate,移除硬编码的batch_size等。

完整适配后的代码

from __future__ import print_function, division
import tensorflow as tf
from tensorflow.keras.datasets import mnist
from tensorflow.keras.layers import Input, Dense, Reshape, Flatten, Dropout
from tensorflow.keras.layers import BatchNormalization, Activation, ZeroPadding2D
from tensorflow.keras.layers import LeakyReLU
from tensorflow.keras.layers import Conv2D, UpSampling2D
from tensorflow.keras.models import Sequential
from tensorflow.keras.optimizers import RMSprop
import matplotlib.pyplot as plt
import numpy as np
import os

class RandomWeightedAverage(tf.keras.layers.Layer):
    """Provides a (random) weighted average between real and generated image samples"""
    def __init__(self, batch_size=32):
        super().__init__()
        self.batch_size = batch_size

    def call(self, inputs, **kwargs):
        alpha = tf.random.uniform((self.batch_size, 1, 1, 1))
        return (alpha * inputs[0]) + ((1 - alpha) * inputs[1])

    def compute_output_shape(self, input_shape):
        return input_shape[0]

class WGANGP():
    def __init__(self, height=128, width=128, channels=3, noise_dim=100, batch_size=64):
        self.img_height = height
        self.img_width = width
        self.channels = channels
        self.img_shape = (self.img_height, self.img_width, self.channels)
        self.noise_dim = noise_dim
        self.batch_size = batch_size

        # Following parameter and optimizer set as recommended in paper
        self.n_critic = 5
        # 适配TF2 API,lr改为learning_rate
        self.critic_optimizer = RMSprop(learning_rate=0.00005)
        self.generator_optimizer = RMSprop(learning_rate=0.00005)

        # Build the generator and critic
        self.generator = self.build_generator()
        self.critic = self.build_critic()

        # 初始化插值层实例
        self.random_weighted_average = RandomWeightedAverage(batch_size=self.batch_size)

    def wasserstein_loss(self, y_true, y_pred):
        return tf.reduce_mean(y_true * y_pred)

    def compute_gradient_penalty(self, real_img, fake_img):
        # 生成插值图像
        interpolated_img = self.random_weighted_average([real_img, fake_img])

        # 用GradientTape记录梯度计算过程
        with tf.GradientTape() as tape:
            tape.watch(interpolated_img)
            validity_interpolated = self.critic(interpolated_img, training=True)

        # 计算梯度并计算L2范数
        gradients = tape.gradient(validity_interpolated, interpolated_img)
        gradients_norm = tf.sqrt(tf.reduce_sum(tf.square(gradients), axis=[1,2,3]))
        # 计算梯度惩罚项
        gradient_penalty = tf.reduce_mean(tf.square(1 - gradients_norm))
        return gradient_penalty

    def build_generator(self):
        model = Sequential()

        model.add(Dense(128 * 7 * 7, activation="relu", input_dim=self.noise_dim))
        model.add(Reshape((7, 7, 128)))
        model.add(UpSampling2D())
        model.add(Conv2D(128, kernel_size=4, padding="same"))
        model.add(BatchNormalization(momentum=0.8))
        model.add(Activation("relu"))
        model.add(UpSampling2D())
        model.add(Conv2D(64, kernel_size=4, padding="same"))
        model.add(BatchNormalization(momentum=0.8))
        model.add(Activation("relu"))
        model.add(Conv2D(self.channels, kernel_size=4, padding="same"))
        model.add(Activation("tanh"))

        model.summary()
        return model

    def build_critic(self):
        model = Sequential()

        model.add(Conv2D(16, kernel_size=3, strides=2, input_shape=self.img_shape, padding="same"))
        model.add(LeakyReLU(alpha=0.2))
        model.add(Dropout(0.25))
        model.add(Conv2D(32, kernel_size=3, strides=2, padding="same"))
        model.add(ZeroPadding2D(padding=((0,1),(0,1))))
        model.add(BatchNormalization(momentum=0.8))
        model.add(LeakyReLU(alpha=0.2))
        model.add(Dropout(0.25))
        model.add(Conv2D(64, kernel_size=3, strides=2, padding="same"))
        model.add(BatchNormalization(momentum=0.8))
        model.add(LeakyReLU(alpha=0.2))
        model.add(Dropout(0.25))
        model.add(Conv2D(128, kernel_size=3, strides=1, padding="same"))
        model.add(BatchNormalization(momentum=0.8))
        model.add(LeakyReLU(alpha=0.2))
        model.add(Dropout(0.25))
        model.add(Flatten())
        model.add(Dense(1))

        model.summary()
        return model

    def train(self, epochs, batch_size, sample_interval=50):
        # Load the dataset
        (X_train, _), (_, _) = mnist.load_data()

        # Rescale -1 to 1
        X_train = (X_train.astype(np.float32) - 127.5) / 127.5
        X_train = np.expand_dims(X_train, axis=3)

        # Adversarial ground truths
        valid = -np.ones((batch_size, 1))
        fake = np.ones((batch_size, 1))

        for epoch in range(epochs):
            for _ in range(self.n_critic):
                # ---------------------
                #  Train Discriminator
                # ---------------------
                # Select a random batch of images
                idx = np.random.randint(0, X_train.shape[0], batch_size)
                imgs = X_train[idx]
                # Sample generator input
                noise = np.random.normal(0, 1, (batch_size, self.noise_dim))

                # 用GradientTape记录critic的梯度
                with tf.GradientTape() as tape:
                    fake_img = self.generator(noise, training=True)
                    real_validity = self.critic(imgs, training=True)
                    fake_validity = self.critic(fake_img, training=True)
                    gp = self.compute_gradient_penalty(imgs, fake_img)
                    # 计算critic总损失:wasserstein损失 + 10倍梯度惩罚
                    d_loss = -self.wasserstein_loss(valid, real_validity) + self.wasserstein_loss(fake, fake_validity) + 10 * gp

                # 应用梯度更新critic
                critic_gradients = tape.gradient(d_loss, self.critic.trainable_variables)
                self.critic_optimizer.apply_gradients(zip(critic_gradients, self.critic.trainable_variables))

            # ---------------------
            #  Train Generator
            # ---------------------
            noise = np.random.normal(0, 1, (batch_size, self.noise_dim))
            with tf.GradientTape() as tape:
                generated_img = self.generator(noise, training=True)
                validity = self.critic(generated_img, training=True)
                # 生成器目标是最大化wasserstein距离,所以损失取负
                g_loss = -self.wasserstein_loss(valid, validity)

            # 应用梯度更新生成器
            generator_gradients = tape.gradient(g_loss, self.generator.trainable_variables)
            self.generator_optimizer.apply_gradients(zip(generator_gradients, self.generator.trainable_variables))

            # Plot the progress
            print ("%d [D loss: %f] [G loss: %f]" % (epoch, d_loss.numpy(), g_loss.numpy()))

            # If at save interval => save generated image samples
            if epoch % sample_interval == 0:
                self.sample_images(epoch)

    def sample_images(self, epoch):
        r, c = 5, 5
        noise = np.random.normal(0, 1, (r * c, self.noise_dim))
        gen_imgs = self.generator.predict(noise, verbose=0)

        # Rescale images 0 - 1
        gen_imgs = 0.5 * gen_imgs + 0.5

        fig, axs = plt.subplots(r, c)
        cnt = 0
        for i in range(r):
            for j in range(c):
                axs[i,j].imshow(gen_imgs[cnt, :,:,0], cmap='gray')
                axs[i,j].axis('off')
                cnt += 1
        # 确保images目录存在
        os.makedirs("images", exist_ok=True)
        fig.savefig("images/mnist_%d.png" % epoch)
        plt.close()

if __name__ == '__main__':
    img_width = 28
    img_height = 28
    channels = 1
    wgan = WGANGP(height=img_height, width=img_width, channels=channels, batch_size=32)
    wgan.train(epochs=30000, batch_size=32, sample_interval=100)

额外优化点

  • 在sample_images方法中添加了os.makedirs确保保存目录存在,避免报错。
  • 生成器predict时添加verbose=0减少输出冗余。
  • 移除了不必要的导入(如Concatenate、_Merge等),精简代码。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 17:57:42