4.4 生成对抗网络:两个网络的博弈排程 本节摘要:GAN 是唯一内置"两个模型互为损失"的家族:生成器从随机噪声造样本,判别器鉴定真伪,两者在交替训练中互相逼高。本节讲清博弈目标函数与交替训练的调度节奏,在 CIFAR-10 子集上跑通一个最小 GAN,拆解三种典型不稳定现象(模式崩塌、训练震荡、梯度消失)的成因与缓解手段。GAN 的训练循环形状特殊——3.3 节说过 fit 装不下它,本节正好用 5.1 节预告的手写循环先把骨架立起来。 本节能力清单 阅读完本节,你应当能够: 写出 GAN 双方的目标函数,解释"互为损失"的含义; 按交替节奏组织一轮训练:先判别器、再生成器; 用子类化模型加 GradientTape 写出最小 GAN 训练步;
本节摘要:GAN 是唯一内置"两个模型互为损失"的家族:生成器从随机噪声造样本,判别器鉴定真伪,两者在交替训练中互相逼高。本节讲清博弈目标函数与交替训练的调度节奏,在 CIFAR-10 子集上跑通一个最小 GAN,拆解三种典型不稳定现象(模式崩塌、训练震荡、梯度消失)的成因与缓解手段。GAN 的训练循环形状特殊——3.3 节说过 fit 装不下它,本节正好用 5.1 节预告的手写循环先把骨架立起来。
阅读完本节,你应当能够:
把两个网络拟人化最省笔墨。生成器 G 的目标:把随机噪声 z 变换得让判别器信以为真。判别器 D 的目标:给真图打高分、给 G 的产出打低分。G 的损失是"D 对假图的判决"(越被判真越好),D 的损失是"真图判假加假图判真"的合计。关键在于G 的梯度必须穿过 D 传回来——D 不只是裁判,还是 G 的误差信号发生器:D 越强,它指出的"哪里假"越精细,G 的改进方向越明确。双方能力同步上涨,最终 D 判不动了(真假各半),G 的产出逼近真实分布。

标准训练步的顺序是:先用一批真图与一批假图训练 D(此时 G 冻结),再用一批假图训练 G(此时 D 冻结,G 的梯度穿过 D 回传)。两个网络在同一个批次里各更新一次,且互相冻结——这种"步内交替"的形状正是 fit 装不下的原因,必须手写训练循环(5.1 节正式展开,本节先给骨架):
import tensorflow as tf import numpy as np (xtr, _), _ = tf.keras.datasets.cifar10.load_data() xtr = (xtr[:6000].astype("float32") - 127.5) / 127.5 # 归一到 -1 到 1 BATCH = 64 ds = tf.data.Dataset.from_tensor_slices(xtr).shuffle(6000).batch(BATCH).prefetch(tf.data.AUTOTUNE) def gen_block(): """生成器:噪声 100 维 反卷积到 32 乘 32 乘 3""" return tf.keras.Sequential([ tf.keras.layers.Input(shape=(100,)), tf.keras.layers.Dense(4 * 4 * 128, activation="relu"), tf.keras.layers.Reshape((4, 4, 128)), tf.keras.layers.Conv2DTranspose(64, 4, strides=2, padding="same", activation="relu"), # 8 tf.keras.layers.Conv2DTranspose(32, 4, strides=2, padding="same", activation="relu"), # 16 tf.keras.layers.Conv2DTranspose(3, 4, strides=2, padding="same", activation="tanh"), # 32 ]) def disc_block(): """判别器:图出真伪分数""" return tf.keras.Sequential([ tf.keras.layers.Input(shape=(32, 32, 3)), tf.keras.layers.Conv2D(32, 4, strides=2, padding="same", activation="relu"), # 16 tf.keras.layers.Conv2D(64, 4, strides=2, padding="same", activation="relu"), # 8 tf.keras.layers.Flatten(), tf.keras.layers.Dense(1), # 输出 logits,配 from_logits=True 损失 ]) G, D = gen_block(), disc_block() bce = tf.keras.losses.BinaryCrossentropy(from_logits=True) opt_g = tf.optimizers.Adam(2e-4) opt_d = tf.optimizers.Adam(2e-4) print("blocks built") # 输出:blocks built # 两个优化器分开建——两个网络的更新节奏不同,状态各自记账
注意实现细节:末层激活 tanh 对应输入归一到 -1 到 1 的对称区间;判别器输出 logits,损失用 from_logits=True(3.5 节的配对链在 GAN 里同样适用)。
手写训练步把交替节奏写实。技巧在"求导对象的切换":训练 D 时 tape 只记录 D 的变量;训练 G 时 tape 记录 G 的变量,但前向计算要让 D 参与并冻结其更新:
@tf.function def train_step(real_imgs): noise = tf.random.normal([tf.shape(real_imgs)[0], 100]) # ---- 第一拍:更新判别器 ---- with tf.GradientTape() as tape_d: fake = G(noise, training=False) # G 只出图不更新 d_real = D(real_imgs, training=True) d_fake = D(fake, training=True) # 真 图判 1,假 图判 0 loss_d = bce(tf.ones_like(d_real), d_real) + \ bce(tf.zeros_like(d_fake), d_fake) grads_d = tape_d.gradient(loss_d, D.trainable_variables) opt_d.apply_gradients(zip(grads_d, D.trainable_variables)) # ---- 第二拍:更新生成器 ---- noise = tf.random.normal([BATCH, 100]) with tf.GradientTape() as tape_g: fake = G(noise, training=True) d_fake = D(fake, training=False) # D 只出分不更新 loss_g = bce(tf.ones_like(d_fake), d_fake) # 骗到 1 才算赢 grads_g = tape_g.gradient(loss_g, G.trainable_variables) opt_g.apply_gradients(zip(grads_g, G.trainable_variables)) return loss_d, loss_g d_losses, g_losses = [], [] for epoch in range(20): for batch in ds: ld, lg = train_step(batch) d_losses.append(float(ld)); g_losses.append(float(lg)) print(f"epoch 20: D loss {d_losses[-1]:.3f}, G loss {g_losses[-1]:.3f}") # 参考输出:epoch 20: D loss 1.082, G loss 1.351 # 双方损失在 1 附近缠斗、谁也不归零——这正是博弈的常态: # 归零意味着某方碾压,反而是病态信号
两个冻结动作的实现值得盯住:训练 D 时 G 以 training=False 参与前向(BatchNorm 类层走推理模式)但不在 D 的梯度名单里;训练 G 时 D 以 training=False 参与前向且梯度名单里只有 G——D 被求导但不被更新,G 的信号就穿它回来了。
GAN 训练的病有经典三象。模式崩塌:G 发现某一种产出能骗过 D,就只产那一种,生成多样性消失——判别依据是同一批噪声产出的图高度雷同;缓解手段包括给 D 的输入加噪声、用 minibatch 判别、或换更弱的 G 更新频率。训练震荡:双方损失反复大幅摆动不收敛——学习率各降一档、或给 D 与 G 配不同的更新次数比。D 过强:D 轻松识破一切假图,G 的梯度接近零、学不动——单侧标签平滑(把真图标签从 1.0 改成 0.9)是三行代码的实用软化。
# 标签平滑的最小改动:真图标签 1.0 改 0.9 with tf.GradientTape() as tape_d: d_real = D(real_imgs, training=True) loss_d = bce(tf.fill(tf.shape(d_real), 0.9), d_real) + \ bce(tf.zeros_like(d_fake), d_fake) print("label smoothing applied") # 输出:label smoothing applied # D 的自信被压低,G 拿到的梯度信号更耐用
选它的场景:需要高保真图像生成、图像到图像的翻译、超分辨率重建——"以假乱真"恰是目标时。不选的场景:只需要数据增强或异常检测,自编码器(4.3)更稳更省;只需要表示学习,对比学习类方法更简单。GAN 的工程成本集中在训练不稳定的调试上,启动一个 GAN 项目前先掂量调试预算。
⚠️ 常见坑:把 D 与 G 放进同一个优化器。两个网络的更新节奏、状态记账必须分开——共享优化器是"训练莫名发散"的高频肇因。
💡 关键直觉:GAN 的损失曲线没有"收敛"的传统语义,双方在 1 附近长期缠斗才是健康态。盯曲线不如盯产出样本的多样性与清晰度。
下节看改写时序调度规则的新范式:注意力。