TensorFlow 2.10.1中CVAE小批量训练广播形状不兼容错误
问题描述
使用TensorFlow 2.10.1训练卷积变分自编码器(CVAE),训练数据按批量大小32分组,大部分批量形状为[32, 40, 40, 1],最后一个批量形状为[22, 40, 40, 1]。执行train_step函数处理该小批量时触发形状不兼容错误,但单独调用reparameterize函数处理该小批量数据时可正常运行,输出形状为(22, 25)。
错误日志
File "C:\Users\..\AppData\Local\Temp\ipykernel_3284\424814309.py", line 79, in compute_loss z_train = model.reparameterize(mean, logvar) File "C:\Users\..\AppData\Local\Temp\ipykernel_3284\424814309.py", line 56, in reparameterize return eps * tf.exp(logvar) + mean Node: 'mul' required broadcastable shapes [[{{node mul}}]] [Op:__inference_train_step_34072]
模型与训练代码
模型定义
import tensorflow as tf from keras.layers import ELU, PReLU, LeakyReLU from keras.layers import Conv2D, Conv2DTranspose, Dense, Input, Reshape, Flatten from keras.models import Model import numpy as np from keras import backend as K tf.random.set_seed(1234) eps0 = tf.random.normal(shape=(3000, z_dim), seed=66) # make model deterministic class CVAE(tf.keras.Model): """Convolutional variational autoencoder.""" def __init__(self, z_dim, input_shape=(40, 40, 1), deterministic=False): super(CVAE, self).__init__() self.deterministic = deterministic self.z_dim = z_dim weight_init = tf.keras.initializers.GlorotUniform(seed=42) # define encoder encoder_input = Input(shape=input_shape) en_conv1 = Conv2D(filters = 32, kernel_size=4, strides=2, padding='same', kernel_initializer=weight_init)(encoder_input) en_conv1 = LeakyReLU(0.1)(en_conv1) en_conv2 = Conv2D(filters = 64, kernel_size=4, strides=2, padding='same', kernel_initializer=weight_init)(en_conv1) en_conv2 = LeakyReLU(0.1)(en_conv2) en_conv3 = Conv2D(filters = 128, kernel_size=4, strides=2, padding='same', kernel_initializer=weight_init)(en_conv2) en_conv3 = LeakyReLU(0.1)(en_conv3) #en_conv4 = Conv2D(filters = 256, kernel_size=4, strides=2, padding='same', kernel_initializer=weight_init)(en_conv3) #en_conv4 = LeakyReLU(0.1)(en_conv4) en_fc1= Flatten()(en_conv3) mean = Dense(z_dim, kernel_initializer=weight_init)(en_fc1) logvar = Dense(z_dim, kernel_initializer=weight_init)(en_fc1) self.encoder = Model(encoder_input,[mean, logvar], name='encoder') # define decoder decoder_input = Input(shape=(z_dim,)) de_fc2 = Dense(en_conv3.shape[1]*en_conv3.shape[2]*en_conv3.shape[3], activation='relu', kernel_initializer=weight_init)(decoder_input) de_fc2 = Reshape((en_conv3.shape[1], en_conv3.shape[2],en_conv3.shape[3]))(de_fc2) de_conv1 = Conv2DTranspose(filters=128, kernel_size=4, strides=2, activation='relu', padding='same', kernel_initializer=weight_init)(de_fc2) de_conv2 = Conv2DTranspose(filters=64, kernel_size=4, strides=2, activation='relu', padding='same', kernel_initializer=weight_init)(de_conv1) #de_conv3 = Conv2DTranspose(filters=32, kernel_size=4, strides=2, activation='relu', padding='same')(de_conv2) decoder_output = Conv2DTranspose(filters=1, kernel_size=4, strides=2, activation='sigmoid', padding='same', kernel_initializer=weight_init)(de_conv2) self.decoder = Model(decoder_input, decoder_output, name='decoder') @tf.function(reduce_retracing=True) def encode(self, x): mean, logvar = self.encoder(x) return mean, logvar def reparameterize(self, mean, logvar): # sample latent if self.deterministic: eps = eps0[:mean.shape[0], :] else: eps = tf.random.normal(shape=mean.shape) return eps * tf.exp(logvar) + mean def decode(self, z): reconstr = self.decoder(z) return reconstr def compute_kernel(x, y): x_size = K.shape(x)[0] y_size = K.shape(y)[0] dim = K.shape(x)[1] tiled_x = K.tile(K.reshape(x, [x_size, 1, dim]), [1, y_size, 1]) tiled_y = K.tile(K.reshape(y, [1, y_size, dim]), [x_size, 1, 1]) return K.exp(-K.mean(K.square(tiled_x - tiled_y), axis=2) / K.cast(dim, 'float32')) def compute_mmd(x, y): x_kernel = compute_kernel(x, x) y_kernel = compute_kernel(y, y) xy_kernel = compute_kernel(x, y) return K.mean(x_kernel) + K.mean(y_kernel) - 2 * K.mean(xy_kernel) def compute_loss(model, x, deterministic=False): mean, logvar = model.encode(x) z_train = model.reparameterize(mean, logvar) # q(z|x) if deterministic: sample_latent = eps0[-200:, :] else: sample_latent = tf.random.normal(tf.stack([200, z_dim])) # p(z) x_reconstr = model.decode(z_train) loss_mmd = compute_mmd(sample_latent, z_train) loss_rec = K.mean(K.square(x - x_reconstr)) return loss_mmd, loss_rec @tf.function(reduce_retracing=True) def train_step(model, x, optimizer, deterministic=False): """Executes one training step and returns the loss. This function computes the loss and gradients, and uses the latter to update the model's parameters. """ with tf.GradientTape() as tape: loss_mmd, loss_nll = compute_loss(model, x, deterministic) loss = loss_mmd + loss_nll gradients = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables)) return {'loss_mmd': loss_mmd, 'loss_nll': loss_nll, 'loss': loss}
训练代码
batch_size = 32 z_dim = 5 optimizer = tf.keras.optimizers.Adam(1e-4) deterministic = True model = CVAE(z_dim, input_shape=(40, 40, 1), deterministic=deterministic) data_train=np.random.rand(2166, 40, 40, 1) train_dataset = tf.data.Dataset.from_tensor_slices(tf.convert_to_tensor(data_train, dtype=tf.float32)).batch(batch_size) for train_x in train_dataset: d = train_step(model, train_x, optimizer, deterministic) # <---- works fine if train_x is [32,40,40,1], not if [22,40,40,1]
问题原因与解决方案
原因
问题出在reparameterize方法中:
- 当
deterministic=True时,eps是从预定义的eps0切片得到的张量,但在tf.function图模式下,使用mean.shape[0]获取的是静态形状,无法适配动态变化的小批量大小(如22),导致eps的形状与logvar的形状无法广播。 - 非图模式下Python会自动适配动态形状,但图模式下TensorFlow依赖静态形状推导,引发不兼容错误。
解决方法
- 修改
reparameterize方法,使用tf.shape(mean)获取动态批量大小,并将方法装饰为tf.function确保图模式兼容性:
@tf.function(reduce_retracing=True) def reparameterize(self, mean, logvar): # sample latent if self.deterministic: # 使用tf.shape获取动态批量大小 batch_size = tf.shape(mean)[0] eps = self.eps0[:batch_size, :] else: eps = tf.random.normal(shape=tf.shape(mean)) return eps * tf.exp(logvar) + mean
- 将
eps0移至CVAE类的初始化方法中,避免全局变量依赖:
def __init__(self, z_dim, input_shape=(40, 40, 1), deterministic=False): super(CVAE, self).__init__() self.deterministic = deterministic self.z_dim = z_dim # 将eps0作为类属性初始化 self.eps0 = tf.random.normal(shape=(3000, z_dim), seed=66) weight_init = tf.keras.initializers.GlorotUniform(seed=42) # ... 其余初始化代码不变
- 可选:在
compute_loss中,sample_latent的获取也建议使用动态形状,避免依赖全局变量:
if deterministic: sample_latent = model.eps0[-200:, :] else: sample_latent = tf.random.normal(shape=(200, model.z_dim))
内容的提问来源于stack exchange,提问作者Alessandro Bitetto
相关产品推荐
相关产品推荐

