☰
[论文笔记]自监督sketch-to-image生成:从自动编码器到GAN的Self-Supervised Sketch-to-Image Synthesis实践
2026/10/4 21:59:13 网站建设 项目流程

1. 复现这篇 Self-Supervised Sketch-to-Image 到底难在哪

如果你正在搜 sketch-to-image、Self-Supervised、GAN 复现相关的资料,大概率已经看过那篇 AAAI 2021 的《Self-Supervised Sketch-to-Image Synthesis》。这篇论文的核心思路其实不复杂:用自动编码器把草图和 RGB 图像的内容、风格特征解耦,再用 GAN 去细化高分辨率细节,整个训练过程不需要成对的草图数据。听起来很美好,但真正动手复现的时候,坑比想象中多。

我自己在跑这套代码的时候,最先卡住的不是模型结构,而是环境依赖和数据读取。原仓库用的是比较老的 PyTorch 版本,直接 pip install 最新版会报一堆 API 不兼容的问题。另外作者提供的 edges2shoe 数据集里,sketch 和 image 的下标对应关系有错位,如果不把 frame 打印出来检查,训练 loss 会一直震荡不收敛。这些问题在 issue 区基本没人回复,只能自己啃。

这篇文章的目标很明确:帮你把环境配置、模型训练参数、推理脚本全部跑通,并且给出生成质量的对比验证步骤。适合谁看?有一定 PyTorch 基础、想复现 GAN 类论文、或者正在做 sketch-to-image 相关实验的同学。如果你只是想在线体验一下效果,作者给的 Playform 演示需要注册还要 credit,不太划算,不如本地跑。

整个流程我会拆成六块:先讲清楚原问题和场景,再说明为什么需要一个统一的 API 通道来管理实验,然后给出可复制的配置片段,接着验证请求是否成功,再列出常见的报错和排查方法,最后给出后续实验的入口。全程不涉及任何网络工具,所有请求都走合规的 API 通道。

2. 为什么实验环境要统一走 TaoToken API 通道

做论文复现的时候,最烦的事情之一就是环境碎片化。你可能在本地跑训练,在服务器上跑推理,又想用某个在线模型做风格转移的对比实验。如果每个环节都单独配一套 key 和 endpoint,管理起来非常乱。我试过把实验相关的模型调用统一到一个 API 通道上,这样无论是本地脚本还是远程 notebook,都只需要维护一份配置。

TaoToken 在这里的角色就是一个统一的 API 入口。它的官网是 https://taotoken.net/?utm_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite&utm_content= ,API 地址是 https://taotoken.net/api 。注意 API 地址后面不加 UTM 参数,直接请求即可。对于 sketch-to-image 这种需要反复调用模型做对比的实验来说,统一通道的好处是:你不需要在代码里硬编码多个 key,也不需要为每个模型单独写一套请求逻辑。

具体到这篇论文的复现,你可能会用到两类调用:一类是本地训练好的 GAN 模型做推理,另一类是调用在线模型做风格转移的 baseline 对比。前者不涉及 API,后者可以通过 TaoToken 的模型对话接口来完成。比如你想验证生成的草图在风格转移后的效果,可以把生成的图像描述或者特征向量发给模型,让它给出语义一致性的判断。这样你就不需要自己再训一个分类器。

另外,如果你后续要做长期的编码实验或者 Agent 相关的自动化流程,可以考虑 Coding Plan。它适合需要持续调用模型、跑批量任务的场景。对于单次验证模型效果,直接用模型对话就够了。接入文档在 https://taotoken.net/doc ,API Keys 管理在 https://taotoken.net/api-keys 。这些地址都带上了 utm_source=taotoken_aicg_blog_end 和 utm_campaign=rewrite,方便你直接点进去。

需要强调的是,TaoToken 不是用来替代编辑器或者训练框架的,它只是帮你把模型调用这一层统一起来。训练还是在本地或者你自己的服务器上跑,GAN 的优化器、学习率、batch size 这些参数一个都不能少。下面我会给出完整的配置片段。

3. 可复制的环境配置与训练参数

这一节是全文的核心,我会给出可以直接复制粘贴的配置文件。首先是 Python 环境,建议用 conda 创建一个独立环境,Python 版本选 3.8,PyTorch 选 1.7.1 加 CUDA 11.0。原仓库的 requirements 里有些包版本太老,我整理了一份能跑通的版本。

conda create -n s2i python=3.8 conda activate s2i pip install torch==1.7.1+cu110 torchvision==0.8.2+cu110 -f https://download.pytorch.org/whl/torch_stable.html pip install numpy==1.19.5 scipy==1.5.4 pillow==8.0.1 tqdm==4.50.2 tensorboard==2.4.1

接下来是模型训练的参数配置。原论文用了两个阶段的训练:第一阶段训练自动编码器做内容风格解耦,第二阶段用 GAN 做细节细化。我建议把配置写成一个 JSON 文件,方便修改和复现。

{ "dataset": "edges2shoe", "data_root": "./datasets/edges2shoe", "batch_size": 8, "lr_ae": 0.0002, "lr_gan": 0.0001, "beta1": 0.5, "beta2": 0.999, "n_epochs_ae": 50, "n_epochs_gan": 100, "lambda_content": 10.0, "lambda_style": 5.0, "lambda_adv": 1.0, "lambda_dmi": 0.1, "image_size": 256, "sketch_channels": 1, "rgb_channels": 3, "style_dim": 128, "content_dim": 256, "num_sketches_per_image": 4, "checkpoint_dir": "./checkpoints", "log_dir": "./logs" }

这里有几个参数需要特别注意。lambda_dmi是动量互信息最小化损失的权重,原论文里这个值对解耦效果影响很大,设得太小会导致内容和风格混在一起,设得太大又会让生成图像模糊。我实测下来 0.1 比较稳。num_sketches_per_image控制 TOM 模型为每张 RGB 图像生成多少张配对草图,原论文用的是 4,显存不够可以降到 2。

如果你要用 TaoToken 做在线验证,可以在配置里加一段 API 相关的设置。注意 Base URL、Key、Model ID 这三件套要写全。

{ "api_base_url": "https://taotoken.net/api", "api_key": "你的_API_KEY", "model_id": "claude-sonnet-4-20250514", "api_timeout": 30 }

API Key 在 https://taotoken.net/api-keys 这里创建,创建的时候记得选对权限。接入文档里有详细的请求示例,地址是 https://taotoken.net/doc 。如果你用的是 Claude Code 或者类似的编码工具,可以参考 https://taotoken.net/claude-code-anthropic 这个页面里的配置说明。

训练脚本的启动命令如下:

python train_step_1_ae.py --config configs/edges2shoe_ae.json python train_step_2_gan.py --config configs/edges2shoe_gan.json

注意原仓库的train_step_2_gan.py里有一行导入写错了,需要手动改成from evaluate.generate_image_matrix import make_matrix。这个 bug 我提过 issue 但没人回,你自己改一下就行。

4. 验证请求与生成结果检查

训练跑起来之后,怎么确认模型真的在工作?我一般分三步验证。第一步是检查数据加载是否正确,把 dataloader 里的 frame 打印出来,确认 sketch 和 image 的下标是对应的。原数据集在 edges2shoe 上问题不大,但在 art 数据集上错位很严重。

for i, batch in enumerate(train_loader): sketch = batch['sketch'] image = batch['image'] print(f"batch {i}, sketch shape: {sketch.shape}, image shape: {image.shape}") if i == 0: print(f"sketch path: {batch['sketch_path'][0]}") print(f"image path: {batch['image_path'][0]}") if i > 5: break

第二步是检查 loss 曲线。自动编码器阶段的 content loss 和 style loss 应该稳步下降,如果震荡剧烈,多半是学习率太大或者 batch size 太小。GAN 阶段的 adversarial loss 会有波动,但判别器的准确率不应该一直停在 0.5 附近,那说明生成器太弱了。

第三步是生成质量对比。原论文的评估指标有 FID 和 LPIPS,但自己跑的时候不用那么复杂,直接肉眼对比就行。我建议生成一组图像,然后和原图、草图放在一起看。下面是一个推理脚本的示例:

import torch from models.autoencoder import SketchAutoEncoder from models.gan import RefinementGAN from PIL import Image import torchvision.transforms as T def load_model(ae_path, gan_path, device='cuda'): ae = SketchAutoEncoder(style_dim=128, content_dim=256).to(device) ae.load_state_dict(torch.load(ae_path, map_location=device)) ae.eval() gan = RefinementGAN().to(device) gan.load_state_dict(torch.load(gan_path, map_location=device)) gan.eval() return ae, gan def inference(sketch_path, style_image_path, ae, gan, device='cuda'): transform = T.Compose([ T.Resize((256, 256)), T.ToTensor(), T.Normalize(mean=[0.5], std=[0.5]) ]) sketch = transform(Image.open(sketch_path).convert('L')).unsqueeze(0).to(device) style_img = transform(Image.open(style_image_path).convert('RGB')).unsqueeze(0).to(device) with torch.no_grad(): content_feat, style_feat = ae.encode(sketch, style_img) coarse = ae.decode(content_feat, style_feat) refined = gan(coarse, sketch) return coarse, refined ae, gan = load_model('./checkpoints/ae_best.pth', './checkpoints/gan_best.pth') coarse, refined = inference('./samples/sketch_01.png', './samples/style_01.jpg', ae, gan) T.ToPILImage()(refined.squeeze(0).cpu() * 0.5 + 0.5).save('./output/refined_01.png')

跑完这个脚本,你会得到一张细化后的生成图像。如果图像有明显的语义错误,比如鞋子变成了包,那说明内容编码器没学好。如果风格不对,比如颜色和 style image 差很远,那是风格编码器的问题。

如果你想用 TaoToken 的模型对话接口做辅助验证,可以把生成图像的描述发给模型,让它判断语义是否一致。请求示例如下:

curl -X POST https://taotoken.net/api/v1/chat/completions \ -H "Authorization: Bearer 你的_API_KEY" \ -H "Content-Type: application/json" \ -d '{ "model": "claude-sonnet-4-20250514", "messages": [ {"role": "user", "content": "这是一张生成的鞋子图像,请判断它是否符合草图的语义,并给出风格一致性评分。"} ] }'

注意 API 地址是 https://taotoken.net/api ,不要加 UTM 参数。模型对话的入口在 https://taotoken.net/models ,你可以先在那里测试一下请求是否通。

5. 常见报错与排查方法

复现过程中遇到的报错,我整理了几个高频的。第一个是ModuleNotFoundError: No module named 'evaluate',这个是因为原仓库的目录结构有问题,你需要把evaluate文件夹放到和训练脚本同级的目录下,或者在sys.path里加上路径。

第二个是RuntimeError: CUDA out of memory。这个在 batch size 设成 8 的时候很容易出现,尤其是 GAN 阶段。解决办法是把 batch size 降到 4,或者把num_sketches_per_image从 4 降到 2。如果还是不行,就把图像尺寸从 256 降到 128,但这样生成质量会下降。

第三个是ValueError: Expected input batch_size (8) to match target batch_size (4)。这个多半是数据加载器的问题,sketch 和 image 的 batch 对不上。检查一下collate_fn是不是写错了,或者数据集里有没有损坏的图片。

第四个是401 Unauthorized。如果你在调用 TaoToken API 的时候遇到这个,先检查 API Key 是不是复制错了,注意不要有多余的空格。然后确认请求头里的Authorization格式是Bearer 你的_KEY。如果还是不行,去 https://taotoken.net/api-keys 重新生成一个 key。

第五个是local proxy failed。这个报错通常是因为你的环境里设置了代理,但代理不可用。检查一下http_proxy和https_proxy环境变量,如果不需要就 unset 掉。注意我们全程不涉及任何网络工具,所有请求都走直连。

第六个是reading choices相关的错误。这个一般出现在解析 API 响应的时候,说明返回的 JSON 结构和你预期的不一样。建议先把原始响应打印出来看看,确认choices字段是否存在。如果返回的是错误信息,里面会有具体的错误码。

第七个是 OAuth 相关的报错。如果你用的是 Claude Code 或者类似的工具,可能会遇到 OAuth token 过期的问题。解决办法是重新走一遍授权流程,或者直接用 API Key 的方式接入。Claude Code 的配置说明在 https://taotoken.net/claude-code-anthropic 。

最后一个坑是数据集本身的问题。原论文提供的 art 数据集里,sketch 和 image 的对应关系是错的,我调了半天才发现。建议你先把每个样本的路径打印出来,人工检查几组。如果错位严重,就只用 edges2shoe 数据集做实验。

6. 后续实验与统一入口

跑通这篇论文之后,你可以做几个延伸实验。第一个是风格混合,把两张不同风格的 RGB 图像的特征做插值,看生成的草图会是什么样。第二个是风格转移,用一张草图配多张风格图,观察生成结果的多样性。第三个是对比实验,把 TOM 生成的草图和 Canny、HED 的边缘图做对比,看看哪种更适合做 sketch-to-image 的输入。

如果你要做长期的编码实验,或者想把整个流程自动化,可以考虑 Coding Plan。它适合需要持续调用模型、跑批量任务的场景。入口在 https://taotoken.net/coding-plan 。对于单次验证模型效果,直接用模型对话就够了,地址是 https://taotoken.net/models 。

API Keys 的管理在 https://taotoken.net/api-keys ,接入文档在 https://taotoken.net/doc 。官网首页是 https://taotoken.net/?utm_source=taotoken_aicg_blog_end&utm_medium=csdn&utm_campaign=rewrite&utm_content= 。这些地址都带上了归因参数,方便你直接访问。

最后说一个实用技巧:训练 GAN 的时候,把生成器的输出每隔几个 epoch 保存一次,这样你可以看到图像从模糊到清晰的过程。如果中间某个 epoch 突然崩了,可以回滚到之前的 checkpoint。另外,判别器的学习率不要设得比生成器高太多,否则生成器很难学到东西。我一般把判别器的 lr 设成生成器的 0.5 倍。

代码仓库的地址是 https://github.com/odegeasslbc/Self-Supervised-Sketch-to-Image-Synthesis-PyTorch ,论文地址是 https://arxiv.org/abs/2012.09290 。数据集我放在百度网盘了,链接和提取码都是 1111。跑通之后你会发现,作者论文里吹的效果确实有水分,但整体思路还是值得学习的。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询