1. 对抗生成网络(GAN)实战:基于PyTorch的鸢尾花数据生成
GAN(Generative Adversarial Network)是深度学习领域最具创造力的模型之一。我在实际项目中经常用它来生成训练数据不足时的补充样本。今天就用PyTorch带大家实现一个能生成鸢尾花数据的GAN模型,这个案例特别适合想入门生成模型的开发者。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 项目环境与数据准备
2.1 环境配置
首先确保安装了必要的库:
bash复制pip install torch numpy pandas scikit-learn matplotlib
我推荐使用Python 3.8+和PyTorch 1.12+版本。如果设备支持CUDA,可以大幅加速训练:
python复制import torch
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
注意:如果没有GPU,建议减少后续的EPOCHS参数值,否则训练时间会很长
2.2 数据加载与预处理
我们使用经典的鸢尾花数据集,但只选择Setosa类别(类别0)作为真实数据:
python复制from sklearn.datasets import load_iris
from sklearn.preprocessing import MinMaxScaler
iris = load_iris()
X = iris.data
y = iris.target
X_class0 = X[y == 0] # 只选择Setosa类别
# 归一化到[-1,1]范围
scaler = MinMaxScaler(feature_range=(-1, 1))
X_scaled = scaler.fit_transform(X_class0)
这里选择MinMaxScaler而不是StandardScaler,因为GAN的生成器通常使用tanh激活函数,其输出范围正好是[-1,1]。
3. GAN模型构建
3.1 生成器设计
生成器的作用是将随机噪声转换为逼真的数据样本:
python复制class Generator(nn.Module):
def __init__(self, latent_dim=10):
super().__init__()
self.model = nn.Sequential(
nn.Linear(latent_dim, 16),
nn.ReLU(),
nn.Linear(16, 32),
nn.ReLU(),
nn.Linear(32, 4), # 输出维度与数据特征维度一致
nn.Tanh()
)
def forward(self, x):
return self.model(x)
关键点说明:
- 输入层维度latent_dim是潜在空间的维度,决定了噪声的复杂度
- 使用ReLU激活函数加速收敛
- 最后使用Tanh将输出限制在[-1,1]范围
3.2 判别器设计
判别器需要区分真实数据和生成数据:
python复制class Discriminator(nn.Module):
def __init__(self):
super().__init__()
self.model = nn.Sequential(
nn.Linear(4, 32),
nn.LeakyReLU(0.2),
nn.Linear(32, 16),
nn.LeakyReLU(0.2),
nn.Linear(16, 1),
nn.Sigmoid()
)
def forward(self, x):
return self.model(x)
设计要点:
- 使用LeakyReLU防止梯度消失,负斜率设为0.2是常见选择
- 最后一层用Sigmoid输出0-1的概率值
- 结构上与生成器对称但方向相反
4. 模型训练过程
4.1 初始化与参数设置
python复制# 超参数设置
LATENT_DIM = 10
EPOCHS = 10000
BATCH_SIZE = 32
LR = 0.0002
BETA1 = 0.5
# 初始化模型
generator = Generator(LATENT_DIM).to(device)
discriminator = Discriminator().to(device)
# 损失函数和优化器
criterion = nn.BCELoss()
g_optimizer = optim.Adam(generator.parameters(), lr=LR, betas=(BETA1, 0.999))
d_optimizer = optim.Adam(discriminator.parameters(), lr=LR, betas=(BETA1, 0.999))
经验分享:Adam优化器的beta1参数设为0.5可以稳定GAN的训练
4.2 训练循环实现
GAN需要交替训练判别器和生成器:
python复制for epoch in range(EPOCHS):
for real_data in dataloader:
real_data = real_data[0].to(device)
batch_size = real_data.size(0)
# 训练判别器
d_optimizer.zero_grad()
# 真实数据损失
real_output = discriminator(real_data)
d_loss_real = criterion(real_output, torch.ones(batch_size, 1).to(device))
# 生成数据损失
noise = torch.randn(batch_size, LATENT_DIM).to(device)
fake_data = generator(noise).detach() # 阻断梯度流向生成器
fake_output = discriminator(fake_data)
d_loss_fake = criterion(fake_output, torch.zeros(batch_size, 1).to(device))
d_loss = d_loss_real + d_loss_fake
d_loss.backward()
d_optimizer.step()
# 训练生成器
g_optimizer.zero_grad()
noise = torch.randn(batch_size, LATENT_DIM).to(device)
fake_data = generator(noise)
fake_output = discriminator(fake_data)
g_loss = criterion(fake_output, torch.ones(batch_size, 1).to(device))
g_loss.backward()
g_optimizer.step()
关键技巧:
- 训练判别器时用detach()阻断生成器梯度
- 生成器的目标是让判别器将假数据判断为真
- 交替训练保持两者能力平衡
5. 结果评估与可视化
5.1 生成新样本
python复制generator.eval()
with torch.no_grad():
noise = torch.randn(50, LATENT_DIM).to(device)
generated_data = generator(noise).cpu().numpy()
# 反归一化
generated_data = scaler.inverse_transform(generated_data)
real_data = scaler.inverse_transform(X_scaled)
5.2 分布对比可视化
python复制fig, axes = plt.subplots(2, 2, figsize=(12, 10))
for i, ax in enumerate(axes.flatten()):
ax.hist(real_data[:, i], bins=10, alpha=0.6, label='Real')
ax.hist(generated_data[:, i], bins=10, alpha=0.6, label='Generated')
ax.set_title(iris.feature_names[i])
ax.legend()
plt.show()
从分布图可以看出,生成的数据在统计特性上与真实数据非常接近。
6. 实战经验与问题排查
6.1 常见问题解决
-
模式崩溃(Mode Collapse)
- 现象:生成器只产生有限的几种样本
- 解决:尝试增加噪声维度、调整学习率、使用Wasserstein GAN
-
判别器过强
- 现象:判别器准确率接近100%
- 解决:降低判别器学习率、减少判别器层数
-
梯度消失
- 现象:损失值不再变化
- 解决:使用LeakyReLU、尝试不同的优化器参数
6.2 调参技巧
- 学习率通常设置在0.0001到0.0005之间
- 批量大小不宜过大,32-128是常见选择
- 潜在空间维度一般取50-200,简单数据可以更小
- 每训练k次判别器后再训练1次生成器(k通常取1-5)
6.3 进阶改进方向
- 使用Conditional GAN实现按类别生成
- 改用WGAN-GP提高训练稳定性
- 添加谱归一化(Spectral Normalization)
- 尝试Progressive Growing技术生成更高维数据
我在实际项目中发现,对于表格数据生成,适当调整网络结构和损失函数可以显著提升生成质量。比如添加特征相关性约束,或者使用自编码器辅助训练。
