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

对抗自编码器无法正常重构数据,判别器Loss趋近0求诊断

对抗自编码器(AAE)训练异常:判别器Loss持续趋近于0

训练时间序列重构任务的对抗自编码器时,出现异常:生成器(Generator)和判别器(Critic)Loss均下降,多轮后判别器Loss趋近于0。预期应为判别器Loss上升、生成器Loss下降,但实际不符。已尝试增加判别器卷积层、调整学习率、修改Batch Size等方案,均无法解决。输入为window_size=100的时间序列数据,隐空间维度20。

网络结构代码

def build_encoder_layer(input_shape, encoder_reshape_shape):
    input_layer = layers.Input(shape=input_shape)

    x = layers.Bidirectional(CUDNNLSTM(units=window_size, return_sequences=True))(input_layer)
    x = layers.Activation(activation='ReLU')(x)  # new part
    x = BatchNormalization()(x)
    x = layers.Dropout(rate=0.2)(x)
   
    # x = layers.Dropout(rate=0.2)(x)

    # new LSTM layer
    x = layers.Bidirectional(CUDNNLSTM(units=100, return_sequences=True))(x)
    x = layers.Activation(activation='ReLU')(x)
    x = BatchNormalization()(x)
    x = layers.Dropout(rate=0.2)(x)
    
    x = layers.Bidirectional(CUDNNLSTM(units=100, return_sequences=True))(x)

    x = layers.Activation(activation='ReLU')(x)
    x = BatchNormalization()(x)
    x = layers.Dropout(rate=0.2)(x)
    
    x = layers.Flatten()(x)
    x = layers.Dense(20)(x)
    x = layers.Reshape(target_shape=encoder_reshape_shape)(x)
    # x = layers.Activation(activation='tanh')(x)
    model = keras.models.Model(input_layer, x, name='encoder')

    return model


def build_generator_layer(input_shape, generator_reshape_shape):
    input_layer = layers.Input(shape=input_shape)

    x = layers.Flatten()(input_layer)
    x = layers.Dense(generator_reshape_shape[0])(x)
    x = layers.Reshape(target_shape=generator_reshape_shape)(x)
    x = layers.Bidirectional(CUDNNLSTM(units=64, return_sequences=True), merge_mode='concat')(x)
    x = layers.Activation(activation='ReLU')(x)  # new part. we need to add dropout after activation for gen
    x = BatchNormalization()(x)
    x = layers.Dropout(rate=0.2)(x)
    x = layers.UpSampling1D(size=2)(x)

    # New LSTM layer
    x = layers.Bidirectional(CUDNNLSTM(units=64, return_sequences=True), merge_mode='concat')(x)
    x = layers.Activation(activation='ReLU')(x)
    x = BatchNormalization()(x)
    x = layers.Dropout(rate=0.2)(x)

    
    # New LSTM layer
    x = layers.Bidirectional(CUDNNLSTM(units=64, return_sequences=True), merge_mode='concat')(x)
    x = layers.Activation(activation='ReLU')(x)
    x = BatchNormalization()(x)
    x = layers.Dropout(rate=0.2)(x)

    x = layers.Bidirectional(CUDNNLSTM(units=64, return_sequences=True), merge_mode='concat')(x)
    x = layers.Activation(activation='ReLU')(x)
    x = BatchNormalization()(x)
    x = layers.Dropout(rate=0.2)(x)
    x = layers.TimeDistributed(layers.Dense(1))(x)
    x = layers.Activation(activation='tanh')(x)  # originally was relu
    model = keras.models.Model(input_layer, x, name='generator')

    return model


def build_critic_x_layer(input_shape):

    input_layer = layers.Input(shape=input_shape)
    x = layers.Conv1D(filters=64, kernel_size=5)(input_layer)
    x = layers.LeakyReLU(alpha=0.2)(x)
    x = layers.Dropout(rate=0.25)(x)
    x = layers.Flatten()(x)
    x = layers.Dense(units=100)(x)
    model = keras.models.Model(input_layer, x, name='critic_x')

    return model


def build_critic_z_layer(input_shape):
    input_layer = layers.Input(shape=input_shape)
    
    x = layers.Conv1D(filters=64, kernel_size=5)(input_layer)
    x = layers.LeakyReLU(alpha=0.2)(x)
    x = layers.Dropout(rate=0.2)(x)
    x = layers.Flatten()(x)
    model = keras.models.Model(input_layer, x, name='critic_z')

    return model

已测试为判别器增加更多卷积层、不同学习率、不同Batch Size,但所有情况下判别器Loss均持续下降。输入数据为window_size=100的时间序列,隐空间维度20。

批量训练代码

Critic X 训练

def critic_x_train_on_batch(x, z):
    # Loss  
    with tf.GradientTape() as tape:
        
        valid_x = critic_x(x)
        x_ = generator(z)
        fake_x = critic_x(x_)

        # Interpolated 
        alpha = tf.random.uniform([batch_size, 1, 1], 0.0, 1.0, dtype=tf.dtypes.float32)
        x_ = tf.cast(x_, dtype='float32')
        x = tf.cast(x, dtype='float32')
        interpolated = alpha * x + (1 - alpha) * x_ 
        
        with tf.GradientTape() as gp_tape:
            gp_tape.watch(interpolated)
            pred = critic_x(interpolated)
        
        grads = gp_tape.gradient(pred, interpolated)
        grad_norm = tf.norm(tf.reshape(grads, (batch_size, -1)), axis=1)
        gp_loss = 10.0*tf.reduce_mean(tf.square(grad_norm - 1.))
#         grads = tf.square(grads)
#         ddx = tf.sqrt(tf.reduce_sum(grads, axis=np.arange(1, len(grads.shape))))
#        gp_loss = tf.reduce_mean((1.0 - ddx) ** 2)
                
        loss1 = wasserstein_loss(-tf.ones_like(valid_x), valid_x)
        loss2 = wasserstein_loss(tf.ones_like(fake_x), fake_x)
        #loss = tf.add_n([loss1, loss2, gp_loss*10.0])        
        loss = loss1 + loss2 + gp_loss
#        loss = tf.reduce_mean(loss)
                        
    gradients = tape.gradient(loss, critic_x.trainable_weights)
    critic_x_optimizer.apply_gradients(zip(gradients, critic_x.trainable_weights))
    return loss

Critic Z 训练

def critic_z_train_on_batch(x, z):
    
    with tf.GradientTape() as tape:

        z_ = encoder(x)   
        valid_z = critic_z(z)             
        fake_z = critic_z(z_) # <- critic_z 
        
        # Interpolated 
        alpha = tf.random.uniform([batch_size, 1, 1], 0.0, 1.0)
        interpolated = alpha * z + (1 - alpha) * z_ 
                
        with tf.GradientTape() as gp_tape:
            gp_tape.watch(interpolated)
            pred = critic_z(interpolated, training=True)
            
        grads = gp_tape.gradient(pred, interpolated)
        grad_norm = tf.norm(tf.reshape(grads, (batch_size, -1)), axis=1)
        gp_loss = 10.0*tf.reduce_mean(tf.square(grad_norm - 1.))

#         grads = tf.square(grads)
#         ddx = tf.sqrt(tf.reduce_sum(grads, axis=np.arange(1, len(grads.shape))))
#         gp_loss = tf.reduce_mean((1.0 - ddx) ** 2)
        
        loss1 = wasserstein_loss(-tf.ones_like(valid_z), valid_z)
        loss2 = wasserstein_loss(tf.ones_like(fake_z), fake_z) 
        loss = loss1 + loss2 + gp_loss
#        loss = tf.reduce_mean(loss)
        
    gradients = tape.gradient(loss, critic_z.trainable_weights)
    critic_z_optimizer.apply_gradients(zip(gradients, critic_z.trainable_weights))
    return loss

生成器与编码器训练

@tf.function
def enc_gen_train_on_batch(x, z):
    with tf.GradientTape() as enc_tape:
        
        z_gen_ = encoder(x, training=True)
        x_gen_ = generator(z, training=False)        
        x_gen_rec = generator(z_gen_, training=False)
        
        fake_gen_x = critic_x(x_gen_, training=False)
        fake_gen_z = critic_z(z_gen_, training=False)
        
        loss1 = wasserstein_loss(fake_gen_x, -tf.ones_like(fake_gen_x))
        loss2 = wasserstein_loss(fake_gen_z, -tf.ones_like(fake_gen_z))
        loss3 = 10.0*tf.reduce_mean(tf.keras.losses.MSE(x, x_gen_rec))

        enc_loss = loss1 + loss2 + loss3
        # enc_loss = loss3
        
    gradients_encoder = enc_tape.gradient(enc_loss, encoder.trainable_weights)
    encoder_optimizer.apply_gradients(zip(gradients_encoder, encoder.trainable_weights))

    with tf.GradientTape() as gen_tape:
        
        z_gen_ = encoder(x, training=False)
        x_gen_ = generator(z, training=True)        
        x_gen_rec = generator(z_gen_, training=True)
        
        fake_gen_x = critic_x(x_gen_, training=False)
        fake_gen_z = critic_z(z_gen_, training=False)
        
        loss1 = wasserstein_loss(fake_gen_x, -tf.ones_like(fake_gen_x))
        loss2 = wasserstein_loss(fake_gen_z, -tf.ones_like(fake_gen_z))
        loss3 = 10.0*tf.reduce_mean(tf.keras.losses.MSE(x, x_gen_rec))

        gen_loss = loss1 + loss2 + loss3
        # gen_loss = loss3
        
    gradients_generator = gen_tape.gradient(gen_loss, generator.trainable_weights)    
    generator_optimizer.apply_gradients(zip(gradients_generator, generator.trainable_weights))    
    return enc_loss, gen_loss

训练Loss变化

Epoch: 1/30, [Dx loss: 1.098870038986206] [Dz loss: -0.10945891588926315] [E loss: 5.278852462768555] [G loss: 3.7521088123321533]
Epoch: 2/30, [Dx loss: 0.5344677567481995] [Dz loss: -0.13788911700248718] [E loss: 3.6803643703460693] [G loss: 3.213960886001587]
Epoch: 3/30, [Dx loss: 0.34554678201675415] [Dz loss: -0.1143367737531662] [E loss: 3.154308319091797] [G loss: 2.820709228515625]
Epoch: 4/30, [Dx loss: 0.2561565041542053] [Dz loss: -0.1354585587978363] [E loss: 3.0694801807403564] [G loss: 3.1366894245147705]
Epoch: 5/30, [Dx loss: 0.20118674635887146] [Dz loss: -0.1930069476366043] [E loss: 2.9434409141540527] [G loss: 3.0156397819519043]
Epoch: 6/30, [Dx loss: 0.16586175560951233] [Dz loss: -0.2611238360404968] [E loss: 3.078233480453491] [G loss: 2.937089443206787]
Epoch: 7/30, [Dx loss: 0.13812018930912018] [Dz loss: -0.2792683243751526] [E loss: 3.0081939697265625] [G loss: 2.8207926750183105]
Epoch: 8/30, [Dx loss: 0.11749842762947083] [Dz loss: -0.3390710949897766] [E loss: 3.2441465854644775] [G loss: 2.751972198486328]
Epoch: 9/30, [Dx loss: 0.09816166758537292] [Dz loss: -0.4120389521121979] [E loss: 3.4268479347229004] [G loss: 2.742722988128662]

问题分析与解决思路

核心问题判断

判别器Loss持续下降并趋近于0,说明判别器完全主导了对抗过程:它能轻松区分真实样本与生成样本,生成器/编码器未能学到有效的特征映射,导致对抗失衡。

具体原因拆解

  1. 判别器与生成器能力不匹配:编码器/生成器使用了多层双向LSTM,参数量与复杂度远高于判别器的简单卷积结构。判别器过于简单,训练初期就能快速掌握区分样本的规律,后续持续压制生成器。
  2. Wasserstein Loss实现或训练逻辑错误:
    • 检查wasserstein_loss的定义,若实现为标准的tf.reduce_mean(y_true * y_pred),则当前Critic的Loss计算逻辑与WGAN常规逻辑相反——通常Critic要让真实样本输出尽可能大,虚假样本输出尽可能小,Loss应为-E[real] + E[fake]。
    • 生成器的Loss计算逻辑也存在偏差:生成器目标是让Critic无法区分虚假与真实样本,而当前Loss设置希望虚假样本输出趋近于-1,与目标相悖。
  3. 训练步长比例失衡:当前代码中判别器与生成器训练次数为1:1,在WGAN-GP中,通常需要训练判别器3-5次再训练一次生成器,避免判别器过快收敛。
  4. 梯度惩罚(GP)有效性存疑:GP的计算可能存在维度不匹配、梯度范数计算错误等问题;GP权重(当前为10)过高或过低,也会影响判别器训练稳定性。
  5. 隐空间分布匹配问题:真实隐变量z的采样分布若与编码器输出z_的分布差异过大,Critic Z会快速学会区分,导致Loss异常。

解决步骤

  1. 调整模型复杂度匹配:
    • 增强判别器能力:增加卷积层数、增大卷积核数量,或替换为LSTM结构适配时序特性;
    • 适当简化生成器/编码器:减少LSTM层数或单元数,缩小模型能力差距。
  2. 修正Wasserstein Loss逻辑:
    • Critic Loss改为tf.reduce_mean(fake_x) - tf.reduce_mean(valid_x) + gp_loss,符合WGAN目标;
    • 生成器Loss改为-tf.reduce_mean(fake_gen_x),让生成样本的Critic输出尽可能接近真实样本。
  3. 调整训练步长比例:每训练1次生成器/编码器,训练
相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 00:08:21