sjkand/dcgan
0
DCGAN 动漫头像生成系统
基于深度卷积生成对抗网络(DCGAN)的动漫头像智能生成平台,支持训练、推理和可视化完整流程。
📁 项目结构
/root/xm/
├── model.py # DCGAN 模型架构(生成器/判别器,支持64×64和128×128分辨率)
├── train.py # 训练脚本(含数据加载、训练循环、检查点保存)
├── app.py # Streamlit 前端交互界面
├── pp.py # 数据集预处理脚本(人脸检测+模糊过滤+自动缩放)
├── requirements.txt # Python 依赖列表
├── checkpoints/ # 模型检查点目录(含训练历史和最终模型)
├── samples/ # 训练过程样本图片(每5轮保存一次)
├── data1/ # 数据集目录(存放动漫头像图片)
└── .streamlit/ # Streamlit 配置目录🛠️ 环境要求
安装命令:
pip install -r requirements.txt🚀 使用方法
1. 数据集准备
将动漫头像图片放入 data1/ 目录,支持 .jpg, .jpeg, .png, .bmp, .tiff 格式。
推荐数据集:kaggle Anime Faces
可选:预处理过滤无效图片
使用人脸检测过滤非人脸图片,并自动缩放至 64×64:
python pp.py --action move --workers 4--action move: 将无效图片移动到removed/目录(推荐)--action delete: 直接删除无效图片--workers: 并行处理线程数(默认CPU核心数,最大4)
预处理功能:
- ✅ 人脸检测(使用动漫专用 LBP 模型)
- ✅ 高度模糊图片过滤(拉普拉斯方差检测)
- ✅ 自动缩放非标准尺寸图片至 64×64
2. 训练模型
python train.py --image_size 64 --batch_size 128 --num_epochs 200参数说明:
训练输出:
checkpoints/generator_final.pth: 最终生成器权重checkpoints/checkpoint_epoch_xxx.pth: 每15轮保存的完整检查点checkpoints/training_history.json: 训练历史(损失曲线数据)samples/epoch_xxx.png: 每5轮保存的生成样本(64张网格图)
3. 启动前端界面
streamlit run app.py访问 http://localhost:8501 即可使用以下功能:
🎨 头像生成
- 支持单次生成 1-64 张动漫头像
- 可固定随机种子实现结果复现
- 支持下载单张图片或打包下载全部
📈 训练过程可视化
- 实时显示生成器/判别器损失曲线
- D(x) 和 D(G(z)) 分数变化趋势
- 训练配置参数展示
🔬 模型架构展示
- 生成器与判别器网络结构详解
- 关键稳定化策略说明(批归一化、LeakyReLU、Dropout等)
- 训练目标函数数学表达式
📊 训练样本演变
- 查看不同训练阶段的生成效果
- 支持滑动选择轮次对比
🧠 模型架构
生成器(Generator)
输入 100 维随机噪声,通过反卷积逐步上采样至目标分辨率:
z (100) → ConvTranspose2d(100→1024) → BN → ReLU → ... → Tanh → 图像判别器(Discriminator)
输入图像,通过卷积逐步下采样至标量概率输出:
图像 → Conv2d(3→96) → LeakyReLU → Dropout → ... → Sigmoid → [0,1] 概率关键稳定化策略
📊 评估指标
训练过程中监控以下指标:
- G_loss:生成器损失(越小越好)
- D_loss:判别器损失(越小越好)
- D(x):判别器对真实图片的平均输出(接近 0.9 为好)
- D(G(z)):判别器对生成图片的平均输出(接近 0.5 为平衡)
📝 示例脚本
快速训练(64×64分辨率,200轮)
python train.py --image_size 64 --batch_size 128 --num_epochs 200高分辨率训练(128×128)
python train.py --image_size 128 --batch_size 64 --num_epochs 300 --ngf 256 --ndf 128仅启动生成界面(使用预训练模型)
streamlit run app.py📦 检查点格式
检查点文件包含以下内容:
{
'netG_state_dict': dict, # 生成器权重
'netD_state_dict': dict, # 判别器权重(完整检查点)
'config': dict, # 训练配置参数
'G_losses': list, # 生成器损失历史
'D_losses': list, # 判别器损失历史
'D_real_accs': list, # D(x) 历史
'D_fake_accs': list, # D(G(z)) 历史
}📱 系统要求
📄 许可证
本项目仅供学习和研究使用。
技术栈:Python 3.11 · PyTorch 2.x · Streamlit · DCGAN 数据集:Anime Face Dataset 训练设备:GPU(CUDA 加速)
