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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 21:45:54