共计 2387 个字符,预计需要花费 6 分钟才能阅读完成。
为什么 GAN 训练需要规范数据集
刚接触生成对抗网络时,最容易忽视的就是数据质量。我曾用网上爬取的动漫头像训练 DCGAN,结果生成全是扭曲的色块。后来发现主要问题出在:

- 图像尺寸从 200×200 到 800×800 不等
- 背景杂乱且有水印
- 样本只有 800 张(至少需要 1 万 +)
这些会导致判别器过早收敛,生成器学不到有效特征。而现成的 MNIST 虽然规整,但:
- 分辨率太低(28×28)
- 灰度图像缺乏色彩信息
- 类别区分度太高(数字间界限明确)
chapter04 数据集设计思路
理想的 GAN 训练数据集应该:
- 分辨率统一(建议 64×64 或 128×128)
- 主题风格一致(如全部人脸或风景)
- 包含适量噪声增强鲁棒性
这里演示生成 10000 张模拟人脸:
# Python 3.8+ 依赖:numpy, opencv, tqdm
import cv2
import numpy as np
from typing import Tuple
def generate_face(rng: np.random.Generator, size: Tuple[int, int] = (128, 128)) -> np.ndarray:
"""
生成带随机噪声的模拟人脸
Args:
rng: 随机数生成器
size: 输出图像尺寸
Returns:
[H,W,3]格式的 RGB 图像,值域 0 -255
"""
# 基础肤色
skin = rng.integers(100, 200, size=(*size, 3))
# 椭圆人脸轮廓
cv2.ellipse(skin, (size[1]//2, size[0]//2),
(size[1]//3, size[0]//2), 0, 0, 360,
(rng.integers(150,200),)*3, -1)
# 添加高斯噪声
noise = rng.normal(0, 15, skin.shape)
return np.clip(skin + noise, 0, 255).astype(np.uint8)
双框架数据管道实现
TensorFlow Dataset 版
import tensorflow as tf
def build_tf_dataset(batch_size=64) -> tf.data.Dataset:
"""构建 TF 数据管道"""
# 生成器函数
def gen():
rng = np.random.default_rng()
while True:
img = generate_face(rng)
yield img.astype(np.float32) / 255.0 # 归一化
return tf.data.Dataset.from_generator(
gen,
output_signature=tf.TensorSpec(shape=(128,128,3), dtype=tf.float32)
).batch(batch_size).prefetch(2)
PyTorch DataLoader 版
import torch
from torch.utils.data import Dataset, DataLoader
class FaceDataset(Dataset):
def __init__(self, num_samples=10000):
self.rng = np.random.default_rng()
self.num_samples = num_samples
def __len__(self):
return self.num_samples
def __getitem__(self, idx):
img = generate_face(self.rng)
return torch.from_numpy(img).permute(2,0,1).float() / 255.0
# 使用时:loader = DataLoader(FaceDataset(), batch_size=64, num_workers=4)
关键避坑技巧
内存泄漏问题
错误写法:
# 在循环内重复创建 Dataset
for epoch in range(100):
ds = build_tf_dataset() # 每次都会新建迭代器
for batch in ds: # 内存会持续增长
train_step(batch)
正确做法:
ds = build_tf_dataset() # 只创建一次
for epoch in range(100):
for batch in ds: # 复用迭代器
train_step(batch)
多 GPU 数据分片
PyTorch 中需配合 DistributedSampler:
sampler = torch.utils.data.distributed.DistributedSampler(dataset)
loader = DataLoader(dataset, sampler=sampler)
性能优化实测
测试 10,000 张 128×128 图像的读取速度:
| 存储方式 | 单 epoch 耗时 |
|---|---|
| 直接读取 PNG | 32.4s |
| HDF5 单文件存储 | 5.1s |
| TFRecords | 7.8s |
HDF5 的写入示例:
import h5py
with h5py.File('faces.h5', 'w') as f:
dset = f.create_dataset('images', (10000,128,128,3), dtype='uint8')
for i in range(10000):
dset[i] = generate_face(np.random.default_rng())
动手实验
尝试修改 generate_face() 函数中的噪声参数:
- 将高斯噪声标准差从 15 改为 30
- 添加泊松噪声(
noise = rng.poisson(10, skin.shape)) - 观察生成图像质量变化对 GAN 训练的影响
结语
构建高质量数据集是 GAN 训练成功的前提。建议先用小批量数据(如 1000 张)快速验证模型结构,再扩展到完整数据集。遇到模式崩溃(mode collapse)时,可以尝试:
- 增加标签噪声
- 调整判别器的更新频率
- 添加谱归一化等正则化手段
正文完
