ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

[论文笔记]自监督sketch-to-image生成:从自动编码器到GAN的Self-Supervised Sketch-to-Image Synthesis实践

[论文笔记]自监督sketch-to-image生成:从自动编码器到GAN的Self-Supervised Sketch-to-Image Synthesis实践 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_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_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_sourcetaotoken_aicg_blog_end 和 utm_campaignrewrite方便你直接点进去。需要强调的是TaoToken 不是用来替代编辑器或者训练框架的它只是帮你把模型调用这一层统一起来。训练还是在本地或者你自己的服务器上跑GAN 的优化器、学习率、batch size 这些参数一个都不能少。下面我会给出完整的配置片段。3. 可复制的环境配置与训练参数这一节是全文的核心我会给出可以直接复制粘贴的配置文件。首先是 Python 环境建议用 conda 创建一个独立环境Python 版本选 3.8PyTorch 选 1.7.1 加 CUDA 11.0。原仓库的 requirements 里有些包版本太老我整理了一份能跑通的版本。conda create -n s2i python3.8 conda activate s2i pip install torch1.7.1cu110 torchvision0.8.2cu110 -f https://download.pytorch.org/whl/torch_stable.html pip install numpy1.19.5 scipy1.5.4 pillow8.0.1 tqdm4.50.2 tensorboard2.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(fbatch {i}, sketch shape: {sketch.shape}, image shape: {image.shape}) if i 0: print(fsketch path: {batch[sketch_path][0]}) print(fimage 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, devicecuda): ae SketchAutoEncoder(style_dim128, content_dim256).to(device) ae.load_state_dict(torch.load(ae_path, map_locationdevice)) ae.eval() gan RefinementGAN().to(device) gan.load_state_dict(torch.load(gan_path, map_locationdevice)) gan.eval() return ae, gan def inference(sketch_path, style_image_path, ae, gan, devicecuda): 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_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_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。跑通之后你会发现作者论文里吹的效果确实有水分但整体思路还是值得学习的。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进