
1. 从一张草图到成品图Pix2Pix 到底在做什么Pix2Pix 是一种基于条件生成对抗网络cGAN的图像到图像翻译模型它能做的事情很具体给你一张输入图输出一张对应的目标图。比如把建筑立面线稿变成带材质的实景图、把黑白照片上色、把卫星图转成地图样式。它适合谁适合已经会写基础 PyTorch 训练循环、想入门 GAN 图像翻译方向的开发者也适合需要快速验证“成对图像数据能不能训出效果”的算法同学。它和普通 GAN 最大的区别在于“成对”。普通 GAN 只要求生成器骗过判别器而 Pix2Pix 要求输入和输出是配对的比如同一场景的线稿和实拍图。生成器采用 U-Net 结构编码器逐层下采样提取语义解码器逐层上采样恢复分辨率中间用跳跃连接把浅层细节直接送到对应层避免边缘和纹理丢失。判别器采用 PatchGAN它不判断整张图真假而是把图切成若干小块逐块判断真假这样能更好地约束局部纹理一致性。损失函数是 Pix2Pix 的关键设计L L_GAN λ * L_L1。L_GAN 负责让生成图看起来“像真的”L1 损失负责让生成图和目标图在像素级接近λ 通常取 100。只靠 L_GAN 容易产生伪影只靠 L1 会模糊两者结合才稳定。这一篇我会带你从零搭出可运行的训练骨架同时用 TaoToken 的统一 Key 接入 AI 工具辅助生成代码片段和排查报错。整个流程分四步准备配置、写模型与训练循环、启动训练、验证结果。2. TaoToken 前置统一 Key 与 API 通道准备在开始写 Pix2Pix 之前先把 AI 辅助工具接好。TaoToken 提供统一的 API 通道一个 Key 可以调用多种模型适合在写代码、查报错、生成配置时随时调用。官网入口是 https://taotoken.net/?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content API 地址是 https://taotoken.net/api 。你需要先拿到 API Key。进入控制台创建 Key路径是 https://taotoken.net/console?utm_sourcetaotoken_aicg_blog_endutm_contentconsoleutm_campaignrewrite 创建完成后在 API Keys 页面复制地址是 https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_contentapi-keysutm_campaignrewrite 。这个 Key 后面会写进配置文件不要直接硬编码在训练脚本里。如果你只是想先验证模型能不能正常对话可以打开模型对话页面 https://taotoken.net/chat?utm_sourcetaotoken_aicg_blog_endutm_contentchatutm_campaignrewrite 输入“帮我解释 PatchGAN 的 receptive field 怎么算”这类问题确认通道可用。如果你打算长期做编码和 Agent 任务可以了解 Coding Plan https://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_contentcoding-planutm_campaignrewrite 它更适合高频调用场景。接入文档在 https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_contentdocutm_campaignrewrite 遇到参数问题先查这里。注意API Key 属于敏感信息建议放在环境变量或本地配置文件不要提交到 Git 仓库。3. 可复制配置config.toml 与 settings.json 骨架为了让训练脚本和 AI 辅助工具都能读到统一配置我习惯把项目参数拆成两份config.toml管训练超参settings.json管 API 通道。这样换数据集或换模型时不用改代码。先建项目目录mkdir pix2pix-demo cd pix2pix-demo mkdir -p data/edges2shoes checkpoints samplesconfig.toml内容如下覆盖数据路径、图像尺寸、batch size、学习率、λ 权重和训练轮数[data] dataroot ./data/edges2shoes image_size 256 batch_size 4 num_workers 2 [train] epochs 200 lr 0.0002 beta1 0.5 beta2 0.999 lambda_l1 100.0 save_every 10 sample_every 5 [model] in_channels 3 out_channels 3 ngf 64 ndf 64settings.json用来放 TaoToken 通道信息Key 从环境变量读取避免明文{ api_base: https://taotoken.net/api, api_key_env: TAOTOKEN_API_KEY, default_model: claude-sonnet, timeout: 60, max_retries: 3 }设置环境变量export TAOTOKEN_API_KEY你的Key读取配置的 Python 代码import tomllib import json import os with open(config.toml, rb) as f: cfg tomllib.load(f) with open(settings.json, r) as f: settings json.load(f) api_key os.environ.get(settings[api_key_env]) assert api_key, 请先设置 TAOTOKEN_API_KEY这样训练脚本只依赖cfgAI 辅助脚本只依赖settings职责清晰。4. 模型与训练循环U-Net 生成器 PatchGAN 判别器4.1 U-Net 生成器实现U-Net 的核心是下采样、上采样和跳跃连接。下面是一个可直接运行的简化版import torch import torch.nn as nn class UNetDown(nn.Module): def __init__(self, in_ch, out_ch, normalizeTrue, dropout0.0): super().__init__() layers [nn.Conv2d(in_ch, out_ch, 4, 2, 1, biasFalse)] if normalize: layers.append(nn.BatchNorm2d(out_ch)) layers.append(nn.LeakyReLU(0.2, inplaceTrue)) if dropout: layers.append(nn.Dropout(dropout)) self.model nn.Sequential(*layers) def forward(self, x): return self.model(x) class UNetUp(nn.Module): def __init__(self, in_ch, out_ch, dropout0.0): super().__init__() layers [ nn.ConvTranspose2d(in_ch, out_ch, 4, 2, 1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ] if dropout: layers.append(nn.Dropout(dropout)) self.model nn.Sequential(*layers) def forward(self, x, skip): x self.model(x) return torch.cat([x, skip], dim1) class GeneratorUNet(nn.Module): def __init__(self, in_ch3, out_ch3, ngf64): super().__init__() self.d1 UNetDown(in_ch, ngf, normalizeFalse) self.d2 UNetDown(ngf, ngf * 2) self.d3 UNetDown(ngf * 2, ngf * 4) self.d4 UNetDown(ngf * 4, ngf * 8, dropout0.5) self.d5 UNetDown(ngf * 8, ngf * 8, dropout0.5) self.d6 UNetDown(ngf * 8, ngf * 8, dropout0.5) self.d7 UNetDown(ngf * 8, ngf * 8, dropout0.5) self.u1 UNetUp(ngf * 8, ngf * 8, dropout0.5) self.u2 UNetUp(ngf * 16, ngf * 8, dropout0.5) self.u3 UNetUp(ngf * 16, ngf * 8, dropout0.5) self.u4 UNetUp(ngf * 16, ngf * 8) self.u5 UNetUp(ngf * 16, ngf * 4) self.u6 UNetUp(ngf * 8, ngf * 2) self.u7 UNetUp(ngf * 4, ngf) self.final nn.Sequential( nn.ConvTranspose2d(ngf * 2, out_ch, 4, 2, 1), nn.Tanh() ) def forward(self, x): d1 self.d1(x) d2 self.d2(d1) d3 self.d3(d2) d4 self.d4(d3) d5 self.d5(d4) d6 self.d6(d5) d7 self.d7(d6) u1 self.u1(d7, d6) u2 self.u2(u1, d5) u3 self.u3(u2, d4) u4 self.u4(u3, d3) u5 self.u5(u4, d2) u6 self.u6(u5, d1) u7 self.u7(u6, x) return self.final(u7)4.2 PatchGAN 判别器实现PatchGAN 输出的是一个特征图每个位置代表原图一个 patch 的真假概率class DiscriminatorPatchGAN(nn.Module): def __init__(self, in_ch6, ndf64): super().__init__() self.model nn.Sequential( nn.Conv2d(in_ch, ndf, 4, 2, 1), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf, ndf * 2, 4, 2, 1, biasFalse), nn.BatchNorm2d(ndf * 2), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf * 2, ndf * 4, 4, 2, 1, biasFalse), nn.BatchNorm2d(ndf * 4), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf * 4, 1, 4, 1, 1), ) def forward(self, x, y): return self.model(torch.cat([x, y], dim1))4.3 训练循环训练循环里生成器和判别器交替更新注意判别器更新时要把生成图 detachimport torch.optim as optim device cuda if torch.cuda.is_available() else cpu G GeneratorUNet(cfg[model][in_channels], cfg[model][out_channels], cfg[model][ngf]).to(device) D DiscriminatorPatchGAN(cfg[model][in_channels] * 2, cfg[model][ndf]).to(device) criterion_gan nn.BCEWithLogitsLoss() criterion_l1 nn.L1Loss() opt_G optim.Adam(G.parameters(), lrcfg[train][lr], betas(cfg[train][beta1], cfg[train][beta2])) opt_D optim.Adam(D.parameters(), lrcfg[train][lr], betas(cfg[train][beta1], cfg[train][beta2])) for epoch in range(cfg[train][epochs]): for i, (real_A, real_B) in enumerate(dataloader): real_A, real_B real_A.to(device), real_B.to(device) fake_B G(real_A) opt_G.zero_grad() pred_fake D(real_A, fake_B) loss_GAN criterion_gan(pred_fake, torch.ones_like(pred_fake)) loss_L1 criterion_l1(fake_B, real_B) * cfg[train][lambda_l1] loss_G loss_GAN loss_L1 loss_G.backward() opt_G.step() opt_D.zero_grad() pred_real D(real_A, real_B) loss_real criterion_gan(pred_real, torch.ones_like(pred_real)) pred_fake_d D(real_A, fake_B.detach()) loss_fake criterion_gan(pred_fake_d, torch.zeros_like(pred_fake_d)) loss_D (loss_real loss_fake) * 0.5 loss_D.backward() opt_D.step()如果你在写这段代码时不确定某个层的输出尺寸可以把问题丢给 TaoToken 的模型对话通道让它帮你算一遍卷积输出公式比手推快很多。5. 验证请求与成功结果训练启动后先跑一个小规模验证确认前向传播和损失计算正常。用随机张量做一次 dry runx torch.randn(2, 3, 256, 256).to(device) y torch.randn(2, 3, 256, 256).to(device) fake G(x) print(生成器输出形状:, fake.shape) d_out D(x, fake) print(判别器输出形状:, d_out.shape)预期输出生成器输出形状: torch.Size([2, 3, 256, 256]) 判别器输出形状: torch.Size([2, 1, 30, 30])判别器输出 30x30 说明 PatchGAN 把 256x256 的图映射成了 30x30 个 patch 的真假判断符合预期。如果形状不对优先检查卷积的 stride 和 padding。正式训练时观察日志正常情况下loss_D会在 0.3 到 0.7 之间波动loss_G前期下降较快后期趋于平稳。如果loss_D迅速降到接近 0说明判别器太强可以降低判别器学习率或增加生成器更新频率。训练 10 个 epoch 后保存一次样本图import torchvision.utils as vutils G.eval() with torch.no_grad(): sample G(real_A[:4]) grid vutils.make_grid(torch.cat([real_A[:4], sample], dim0), nrow4, normalizeTrue) vutils.save_image(grid, fsamples/epoch_{epoch}.png) G.train()打开samples/epoch_10.png上半部分是输入线稿下半部分是生成图。如果生成图开始出现目标域的纹理和颜色说明训练方向正确。6. 本篇常见错排查6.1 判别器输出形状不对报错通常是BCEWithLogitsLoss的 target 和 input 形状不匹配。原因是 PatchGAN 输出的是特征图不是标量。解决方法是不要对判别器输出做view(-1)直接用torch.ones_like(pred)构造同形状标签。6.2 生成图全灰或全黑多半是 L1 权重过大或生成器最后一层激活函数不对。Pix2Pix 生成器输出范围是 [-1, 1]所以最后一层用Tanh同时数据预处理要把像素归一化到 [-1, 1]。如果用了Sigmoid输出范围变成 [0, 1]和 L1 目标不匹配就会训崩。6.3 CUDA out of memory256x256、batch size 4 在 8GB 显存上通常够用。如果不够先把 batch size 降到 2或者把ngf和ndf从 64 降到 32。不要一上来就上 512x512Pix2Pix 在高分辨率下显存增长很快。6.4 API 调用返回 401检查TAOTOKEN_API_KEY是否设置成功可以用echo $TAOTOKEN_API_KEY确认。如果 Key 正确但仍报错检查请求头里的Authorization格式是否为Bearer key。接入细节参考 https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_contentdocutm_campaignrewrite 。6.5 训练 loss 震荡不收敛先确认数据配对是否正确输入图和目标图必须一一对应。其次检查学习率0.0002 是 Pix2Pix 论文的推荐值如果数据集很小可以降到 0.0001。另外 Adam 的 beta1 设成 0.5 而不是默认的 0.9这是 GAN 训练的常见调整。7. 继续推进从跑通到调优跑通训练骨架只是第一步。接下来可以做的调优方向包括把 L1 换成 Perceptual Loss 提升细节、给判别器加谱归一化稳定训练、用学习率衰减策略让后期更稳。如果你需要长期跑编码和 Agent 任务Coding Plan 的通道更适合高频调用地址是 https://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_contentcoding-planutm_campaignrewrite 。模型对话入口在 https://taotoken.net/chat?utm_sourcetaotoken_aicg_blog_endutm_contentchatutm_campaignrewrite 遇到报错可以直接贴日志让它帮你定位。API Key 管理在 https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_contentapi-keysutm_campaignrewrite 建议给训练辅助脚本单独建一个 Key方便轮换。我自己的习惯是每改一次模型结构就先跑 5 个 epoch 看 loss 曲线和样本图确认没有崩再拉长训练。Pix2Pix 对数据配对质量很敏感如果生成图始终模糊先回头检查数据预处理而不是急着换模型。