一、项目简介

图像风格转换是生成式深度学习的经典应用场景,传统监督学习需要一一配对的数据集,训练成本极高。而CycleGAN(循环一致性生成对抗网络)凭借无监督、非配对图像转换的特性,被广泛应用于图像风格迁移、物体纹理转换、场景渲染等领域。

本项目基于PyTorch手写实现轻量化CycleGAN网络,完成马(Horse)→ 斑马(Zebra)的双向图像风格转换任务。全程基于GPU加速训练,完成100轮完整迭代,最终实现基础的跨物种纹理迁移效果,同时记录训练过程中的效果优劣、存在缺陷以及后续优化方案,适合GAN入门学习与实战复盘。

二、CycleGAN核心原理

CycleGAN 核心解决了无配对图像转换的问题,区别于普通GAN,新增循环一致性约束,避免生成网络出现模式崩塌、转换无效的问题,核心由两大模块、三大损失函数构成:

1. 网络结构组成

  • 双生成器:G_A2B(马转斑马)、G_B2A(斑马转马),采用ResNet残差结构,保证深层网络训练不梯度消失

  • 双判别器:D_A(判别真实马匹图像)、D_B(判别真实斑马图像),采用Patch判别结构,逐区块判别图像真伪,细节表现力更强

2. 核心损失函数

  • GAN对抗损失:让生成图像尽可能贴近真实数据集分布,欺骗判别器

  • Cycle循环一致性损失(L1):转换图像还原原图,保证图像主体结构不丢失(核心约束)

  • Identity恒等损失:输入同源图像不做多余转换,保留图像基础特征

三、项目环境与数据集配置

1. 运行环境

本项目基于Anaconda虚拟环境搭建,全程GPU加速训练,通用适配所有NVIDIA显卡:

  • 编程语言:Python 3.12

  • 深度学习框架:PyTorch(CUDA GPU加速)

  • 核心依赖:torch、torchvision、tqdm、Pillow、glob

  • 训练设备:NVIDIA GPU 加速

2. 数据集结构

采用标准非配对斑马/马数据集,严格遵循CycleGAN官方目录规范,无需一一配对,目录结构如下:

datasets/zebra2horse/
├── trainA/    # 训练集A:马匹图像
├── trainB/    # 训练集B:斑马图像
├── testA/     # 测试集A:待转换马匹图像
└── testB/     # 测试集B:待转换斑马图像

数据集统一预处理为256×256尺寸,包含随机水平翻转、归一化等数据增强操作,提升模型泛化能力。

四、整体代码框架设计

本项目代码分为训练模块、测试模块、一键运行模块三部分,模块化拆分清晰,剔除冗余代码,仅展示核心框架,完整代码可自行拓展优化。

1. 数据加载模块

自定义非配对数据集类,自动读取AB两类图像,实现随机采样、数据预处理,适配CycleGAN无配对训练特性:

  • 自动校验数据集文件,避免空数据集报错

  • 统一图像尺寸、归一化处理,适配模型输入

  • 随机抽取不同类别图像,满足非配对训练需求

    # 1.数据集
    class UnalignedDataset(Dataset):
        def __init__(self): pass
        def __len__(self): pass
        def __getitem__(self): pass

2. 网络模型模块

核心实现残差块、ResNet生成器、Patch判别器三大基础网络:

  • ResnetBlock残差块:堆叠9层残差结构,提取深层图像纹理特征,避免梯度消失

  • ResnetGenerator生成器:下采样提取特征+残差特征重构+上采样还原图像尺寸,完成风格转换

  • PatchDiscriminator判别器:逐区块判别图像真伪,提升细节判别精度

    # 2.网络
    class ResnetBlock(nn.Module): pass
    class ResnetGenerator(nn.Module): pass
    class PatchDiscriminator(nn.Module): pass
    

3. 训练逻辑模块

实现标准CycleGAN训练流程,包含权重初始化、优化器配置、损失计算、模型保存、断点续训功能:

  • 双生成器、双判别器交替训练,固定一方参数训练另一方

  • 融合三类损失函数,加权约束模型训练方向

  • 配置图像池化层,缓存生成图像,稳定训练过程

  • 每轮自动保存最新权重,每10轮保存历史权重,方便效果筛选

    # 3.工具函数
    class ImagePool: pass
    def save_checkpoint(): pass
    
    # 4.训练主逻辑
    def train(opt):
        # 初始化数据集、网络、损失、优化器
        # epoch循环:交替训练生成器、判别器
        # 按轮保存权重
        pass
    
    if __name__=="__main__":
        opt = parse_args()
        train(opt)

4. 测试推理模块

加载训练好的权重,支持双向风格转换,批量推理测试图像,自动保存转换结果:

  • 支持指定权重、指定转换方向(AtoB/BtoA)

  • 张量与图像格式互转,适配可视化输出

  • 批量测试文件夹内所有图像,自动分类保存结果

    # 复用生成器结构
    class ResnetBlock(nn.Module): pass
    class ResnetGenerator(nn.Module): pass
    
    # 推理主逻辑
    def test(opt):
        # 加载权重+选择生成器
        # 遍历测试图片,前向推理、保存结果
        pass
    
    if __name__=="__main__":
        opt = parse_args()
        test(opt)

5. 一键运行模块

封装训练、测试、训测联动三种运行模式,可自由切换,无需重复输入命令,降低操作门槛。

# 超参数配置区
MODE = "train_test"

def train(): # 拼接命令启动train.py
def test():  # 拼接命令启动test.py

if __name__=="__main__":
    # 根据MODE分支执行训练/测试
    pass

五、训练参数与过程监控

1. 核心超参数配置

  • 迭代轮数:100 epoch

  • 批次大小:batch_size=1(适配家用GPU显存)

  • 图像尺寸:256×256

  • 学习率:0.0002

  • 循环损失权重:10.0,恒等损失权重:5.0

  • 权重保存频率:每10轮保存一次断点权重

2. 训练过程状态

  • 训练设备:GPU全速运行,稳定速度约3.5~3.6it/s,大概十多个小时,训练全程无显存溢出、无报错崩溃

  • 损失变化:训练前期生成器损失G值较高,随迭代稳步下降;判别器损失D维持稳定合理区间,模型持续学习纹理转换特征

  • 权重保存:最终生成10组阶段性权重+最新实时权重,覆盖全程训练状态

六、实验效果展示

本项目核心实现马匹图像 → 斑马纹理转换,选取训练中效果最优权重进行测试,整体实现基础风格迁移效果。

七、效果优缺点复盘

1. 优势亮点

  • 整体风格转换效果达标,大部分马匹图像可精准完成斑马纹理渲染,主体特征保留完整

  • 训练过程稳定,无模式崩塌、梯度爆炸问题,权重保存完整,可复现性强

  • 模块化代码简洁清晰,适配性强,可快速迁移到其他风格转换任务

2. 现存缺陷

  • 误染色问题:部分图像中非马匹背景区域会被模型错误渲染斑马纹理,背景干扰明显

  • 边缘细节瑕疵:马匹身体轮廓边缘、四肢交界处转换不自然,存在模糊、纹理错位现象

  • 泛化性有限:对角度特殊、光线复杂的图像转换效果较差,细节还原度不足

八、后续优化方案

针对本次训练存在的问题,后续可从数据集、网络结构、训练策略三方面优化,进一步提升转换精度:

  1. 数据集优化:筛选干净无复杂背景的数据集,增加数据增强(随机裁剪、亮度变换),提升模型抗干扰能力

  2. 网络优化:引入注意力机制,让模型聚焦主体区域,抑制背景误染色;调整残差块数量,适配细节纹理学习

  3. 训练策略优化:调整损失函数权重、动态学习率,避免后期模型过拟合;增加模型早停策略,选取最优迭代轮数

  4. 精度优化:引入混合精度训练,在不损失效果的前提下降低显存占用,提升训练速度

九、总结

本次项目完整实现了基于CycleGAN的马-斑马无监督风格转换,从零完成数据集配置、模型训练、批量测试全流程。100轮迭代训练后,模型具备基础的跨物种纹理转换能力,整体效果符合入门级GAN项目预期。

同时项目也暴露了基础CycleGAN模型的通病:背景区分能力弱、细节渲染不足。本次实战完整记录了训练过程与效果缺陷,为后续模型优化、进阶GAN项目学习积累了实战经验。代码模块化程度高,可直接用于学习、复现与二次开发。

    Logo

    汇聚全球AI编程工具,助力开发者即刻编程。

    更多推荐