使用PyTorch后端的Keras3运行VAE模型时遇RuntimeError
问题描述
尝试复现Keras3的卷积VAE模型,切换到PyTorch后端后训练时出现维度不匹配错误。修改后的代码如下:
import os os.environ['KERAS_BACKEND'] = 'torch' import keras import numpy as np from keras import layers, ops class Sampling(layers.Layer): """Uses (z_mean, z_log_var) to sample z, the vector encoding a digit.""" def __init__(self, name='sampling', **kwargs): super(Sampling, self).__init__(name=name, **kwargs) self.seed_generator = keras.random.SeedGenerator(42) def call(self, inputs): z_mean, z_log_var = inputs batch = ops.shape(z_mean)[0] dim = ops.shape(z_mean)[1] epsilon = keras.random.normal(shape=(batch, dim), seed=self.seed_generator) return z_mean + ops.exp(0.5 * z_log_var) * epsilon class Encoder(keras.Model): """Maps MNIST digits to a triplet (z_mean, z_log_var, z).""" def __init__(self, latent_dim=32, intermediate_dim=64, name='encoder', **kwargs): super().__init__(name=name, **kwargs) self.conv_layer1 = layers.Conv2D(32, 3, activation='relu', strides=2, padding='same') self.conv_layer2 = layers.Conv2D(64, 3, activation='relu', strides=2, padding='same') self.flatten = layers.Flatten() self.dense_proj = layers.Dense(intermediate_dim, activation='relu') self.dense_mean = layers.Dense(latent_dim, name='z_mean') self.dense_log_var = layers.Dense(latent_dim, name='z_log_var') self.sampling = Sampling() def call(self, inputs): x = self.conv_layer1(inputs) x = self.conv_layer2(x) x = self.flatten(x) x = self.dense_proj(x) z_mean = self.dense_mean(x) z_log_var = self.dense_log_var(x) z = self.sampling((z_mean, z_log_var)) return z_mean, z_log_var, z class Decoder(keras.Model): """Converts z, the encoded digit vector, back into a readable digit.""" def __init__(self, original_dim, intermediate_dim=64, name='decoder', **kwargs): super(Decoder, self).__init__(name=name, **kwargs) self.dense_proj = layers.Dense(7 * 7 * 64, activation='relu') self.reshape = layers.Reshape((7, 7, 64)) self.conv_transpose1 = layers.Conv2DTranspose(64, 3, activation='relu', strides=2, padding='same') self.conv_transpose2 = layers.Conv2DTranspose(32, 3, activation='relu', strides=2, padding='same') self.dense_output = layers.Conv2DTranspose(1, 3, activation='sigmoid', padding='same') def call(self, inputs): x = self.dense_proj(inputs) x = self.reshape(x) x = self.conv_transpose1(x) x = self.conv_transpose2(x) return self.dense_output(x) class VAE(keras.Model): """Combines the encoder and decoder into an end-to-end model for training.""" def __init__( self, encoder, decoder, name='vae', **kwargs ): super().__init__(name=name, **kwargs) self.encoder = encoder self.decoder = decoder def call(self, input): z_mean, z_log_var, z = self.encoder(input) reconstructed = self.decoder(z) # Add KL divergence regularization loss. kl_loss = -0.5 * ops.mean(1 + z_log_var - ops.square(z_mean) - ops.exp(z_log_var), axis=1) self.add_loss(kl_loss) return reconstructed if __name__ == '__main__': (x_train, _), (x_test, _) = keras.datasets.mnist.load_data() print(x_train.shape) mnist_digits = np.concatenate([x_train, x_test], axis=0) mnist_digits = np.expand_dims(mnist_digits, -1).astype("float32") / 255 print(mnist_digits.shape) encoder = Encoder(latent_dim=32) decoder = Decoder(original_dim=784) vae = VAE(encoder=encoder, decoder=decoder) vae.compile(optimizer=keras.optimizers.Adam(learning_rate=1e-3), loss=keras.losses.MeanSquaredError()) vae.fit(mnist_digits, mnist_digits, epochs=30, batch_size=128)
训练时抛出错误:
/opt/miniconda3/envs/vae/bin/python /Users/belter/github/VAE/example.py (60000, 28, 28) (70000, 28, 28, 1) /opt/miniconda3/envs/vae/lib/python3.11/site-packages/keras/src/backend/common/backend_utils.py:89: UserWarning: You might experience inconsistencies across backends when calling conv transpose with kernel_size=3, stride=2, dilation_rate=1, padding=same, output_padding=1. warnings.warn( Epoch 1/30 Traceback (most recent call last): File "/Users/belter/github/VAE/example.py", line 106, in <module> vae.fit(mnist_digits, mnist_digits, epochs=30, batch_size=128) File "/opt/miniconda3/envs/vae/lib/python3.11/site-packages/keras/src/utils/traceback_utils.py", line 122, in error_handler raise e.with_traceback(filtered_tb) from None File "/opt/miniconda3/envs/vae/lib/python3.11/site-packages/keras/src/backend/torch/numpy.py", line 1248, in stack return torch.stack(x, dim=axis) ^^^^^^^^^^^^^^^^^^^^^^^^ RuntimeError: stack expects each tensor to be equal size, but got [] at entry 0 and [128] at entry 1
环境:Python 3.11、PyTorch 2.2.2、Keras 3.2.1
错误原因
错误根源是损失维度不匹配:
- 主损失(MSE)计算后是标量(整个batch的平均损失)
- 当前KL损失通过
ops.mean(..., axis=1)计算后,得到形状为[batch_size]的张量(每个样本对应一个损失值) - Keras3在PyTorch后端下,
add_loss尝试将主损失与KL损失堆叠时,因维度不一致导致报错
修复方案
1. 调整KL损失计算
将KL损失改为全局平均,输出标量,与主损失维度一致:
# 原代码 kl_loss = -0.5 * ops.mean(1 + z_log_var - ops.square(z_mean) - ops.exp(z_log_var), axis=1) # 修改后 kl_loss = -0.5 * ops.mean(1 + z_log_var - ops.square(z_mean) - ops.exp(z_log_var))
2. 移除Decoder冗余参数
Decoder类中的original_dim参数未被使用,删除该参数并调整实例化代码:
# 修改Decoder类 class Decoder(keras.Model): """Converts z, the encoded digit vector, back into a readable digit.""" def __init__(self, intermediate_dim=64, name='decoder', **kwargs): super(Decoder, self).__init__(name=name, **kwargs) self.dense_proj = layers.Dense(7 * 7 * 64, activation='relu') self.reshape = layers.Reshape((7, 7, 64)) self.conv_transpose1 = layers.Conv2DTranspose(64, 3, activation='relu', strides=2, padding='same', output_padding=1) self.conv_transpose2 = layers.Conv2DTranspose(32, 3, activation='relu', strides=2, padding='same', output_padding=1) self.dense_output = layers.Conv2DTranspose(1, 3, activation='sigmoid', padding='same') # 实例化Decoder时改为 decoder = Decoder()
3. 消除Conv2DTranspose警告
给转置卷积层添加output_padding=1,确保输出尺寸准确(从7x7到28x28),同时消除后端不一致的警告。
完整修复代码
import os os.environ['KERAS_BACKEND'] = 'torch' import keras import numpy as np from keras import layers, ops class Sampling(layers.Layer): """Uses (z_mean, z_log_var) to sample z, the vector encoding a digit.""" def __init__(self, name='sampling', **kwargs): super(Sampling, self).__init__(name=name, **kwargs) self.seed_generator = keras.random.SeedGenerator(42) def call(self, inputs): z_mean, z_log_var = inputs batch = ops.shape(z_mean)[0] dim = ops.shape(z_mean)[1] epsilon = keras.random.normal(shape=(batch, dim), seed=self.seed_generator) return z_mean + ops.exp(0.5 * z_log_var) * epsilon class Encoder(keras.Model): """Maps MNIST digits to a triplet (z_mean, z_log_var, z).""" def __init__(self, latent_dim=32, intermediate_dim=64, name='encoder', **kwargs): super().__init__(name=name, **kwargs) self.conv_layer1 = layers.Conv2D(32, 3, activation='relu', strides=2, padding='same') self.conv_layer2 = layers.Conv2D(64, 3, activation='relu', strides=2, padding='same') self.flatten = layers.Flatten() self.dense_proj = layers.Dense(intermediate_dim, activation='relu') self.dense_mean = layers.Dense(latent_dim, name='z_mean') self.dense_log_var = layers.Dense(latent_dim, name='z_log_var') self.sampling = Sampling() def call(self, inputs): x = self.conv_layer1(inputs) x = self.conv_layer2(x) x = self.flatten(x) x = self.dense_proj(x) z_mean = self.dense_mean(x) z_log_var = self.dense_log_var(x) z = self.sampling((z_mean, z_log_var)) return z_mean, z_log_var, z class Decoder(keras.Model): """Converts z, the encoded digit vector, back into a readable digit.""" def __init__(self, intermediate_dim=64, name='decoder', **kwargs): super(Decoder, self).__init__(name=name, **kwargs) self.dense_proj = layers.Dense(7 * 7 * 64, activation='relu') self.reshape = layers.Reshape((7, 7, 64)) self.conv_transpose1 = layers.Conv2DTranspose(64, 3, activation='relu', strides=2, padding='same', output_padding=1) self.conv_transpose2 = layers.Conv2DTranspose(32, 3, activation='relu', strides=2, padding='same', output_padding=1) self.dense_output = layers.Conv2DTranspose(1, 3, activation='sigmoid', padding='same') def call(self, inputs): x = self.dense_proj(inputs) x = self.reshape(x) x = self.conv_transpose1(x) x = self.conv_transpose2(x) return self.dense_output(x) class VAE(keras.Model): """Combines the encoder and decoder into an end-to-end model for training.""" def __init__( self, encoder, decoder, name='vae', **kwargs ): super().__init__(name=name, **kwargs) self.encoder = encoder self.decoder = decoder def call(self, input): z_mean, z_log_var, z = self.encoder(input) reconstructed = self.decoder(z) # Add KL divergence regularization loss. kl_loss = -0.5 * ops.mean(1 + z_log_var - ops.square(z_mean) - ops.exp(z_log_var)) self.add_loss(kl_loss) return reconstructed if __name__ == '__main__': (x_train, _), (x_test, _) = keras.datasets.mnist.load_data() print(x_train.shape) mnist_digits = np.concatenate([x_train, x_test], axis=0) mnist_digits = np.expand_dims(mnist_digits, -1).astype("float32") / 255 print(mnist_digits.shape) encoder = Encoder(latent_dim=32) decoder = Decoder() vae = VAE(encoder=encoder, decoder=decoder) vae.compile(optimizer=keras.optimizers.Adam(learning_rate=1e-3), loss=keras.losses.MeanSquaredError()) vae.fit(mnist_digits, mnist_digits, epochs=30, batch_size=128)
内容的提问来源于stack exchange,提问作者Belter
相关产品推荐
相关产品推荐

