ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

TensorFlow深度学习实战(28)——CycleGAN详解与实现:用TaoToken统一Key跑通无配对图像风格迁移

TensorFlow深度学习实战(28)——CycleGAN详解与实现:用TaoToken统一Key跑通无配对图像风格迁移 1. 为什么无配对图像翻译总在工程落地时卡住CycleGAN 解决的是一个很具体的痛点手里只有两堆图一堆苹果一堆橘子或者一堆马一堆斑马它们之间没有任何一一对应的关系但你就是想让模型学会“把苹果画成橘子”。传统 pix2pix 那种监督式图像翻译要求成对数据现实里几乎凑不齐CycleGAN 用两个生成器加两个判别器靠循环一致性把“可逆”这件事约束住才让无配对训练变得可行。但真正动手跑的时候问题往往不在论文理解而在工程细节。我见过太多人卡在几个地方TensorFlow 2.x 里tensorflow_examples的导入路径变了、summer2winter_yosemite数据集下载后目录结构和代码对不上、tf.function装饰的train_step里 persistent tape 用错导致梯度为 None、checkpoint 恢复后优化器状态丢失导致 loss 突然跳变。更隐蔽的是凭据管理——当你同时要调多个模型服务、跑多个实验分支时API Key 散落在各个脚本里改一次环境就要翻半天。这篇就按“能直接复制跑通”的标准来写。我会用 TensorFlow 2.x 搭一套完整的 CycleGAN数据集用苹果↔橘子apple2orange和马↔斑马horse2zebra都覆盖训练脚本给全同时把模型调用凭据统一收到 TaoToken 的 Key 通道里管理。最后用固定随机种子加 FID 指标和视觉样例双重验证迁移效果而不是只看 loss 曲线自我安慰。适合谁看已经会写基础 GAN、想把手里的无配对图像数据真正训出可用风格迁移模型的工程师或者正在做图像增强、域适应、数据合成需要一套可复现 CycleGAN 基线的人。你不需要 GPU 集群单卡 8G 显存就能跑 256×256 的配置只是 epoch 数要拉长。先说清楚一个预期CycleGAN 不是“训 10 个 epoch 就出效果”的模型。苹果↔橘子这种颜色域差异明显的大概 40–60 epoch 能看到稳定迁移马↔斑马涉及纹理结构变化通常要 100 epoch 以上。所以 checkpoint 机制必须做对否则中断一次就前功尽弃。2. TaoToken 统一 Key 管理模型调用凭据的前置准备在讲网络结构之前先把凭据这件事理清楚。CycleGAN 本身是本地训练不依赖外部 API但实际工程里你往往不止跑一个模型可能同时要调视觉理解模型做数据清洗、调文本模型生成实验记录、调另一个服务做 FID 评估的辅助计算。这些调用如果各自维护 Key脚本里就会散落一堆硬编码字符串既不安全也不好切换环境。TaoToken 在这里的角色是统一入口一个 Key 走通多个模型服务的调用通道Base URL 固定模型 ID 按需切换。对 CycleGAN 项目来说最直接的用法是把它作为实验管理侧的调用凭据——比如训练完一个 epoch 后自动调模型对话接口生成训练摘要或者用 coding plan 通道跑实验脚本的辅助代码生成。前置准备分三步。第一步拿到 Key。访问https://taotoken.net/api-keys带 UTM 的完整链接是https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_contentcyclegan_tf在控制台里创建一个新 Key复制出来。注意 Key 只在创建时完整显示一次丢了就重新建。第二步确认 Base URL。API 通道统一用https://taotoken.net/api这个地址不加 UTM 参数直接写进配置。模型 ID 根据你要调的服务填比如做对话摘要用对应的对话模型 ID做代码辅助用 coding 相关的模型 ID。三个要素——Base URL、Key、Model ID——缺一不可后面配置文件里会体现。第三步把 Key 放进环境变量而不是代码里。Linux/macOS 下export TAOTOKEN_API_KEYsk-你的实际key export TAOTOKEN_BASE_URLhttps://taotoken.net/apiWindows PowerShell$env:TAOTOKEN_API_KEYsk-你的实际key $env:TAOTOKEN_BASE_URLhttps://taotoken.net/api这样做的意义在于训练脚本、评估脚本、实验记录脚本都从环境变量读换机器或换 Key 时只改一处。如果你用.env文件配合python-dotenv记得把.env加进.gitignore别把 Key 提交上去。注意TaoToken 是模型调用凭据的统一管理通道不是网络代理工具也不替代本地训练环境。CycleGAN 的权重训练全程在你自己的机器上完成TaoToken 只负责你项目里那些需要调外部模型能力的环节。如果你打算长期跑多个实验分支建议直接上 Coding Plan把实验脚本生成、训练日志分析、FID 评估辅助这些环节都挂到统一通道下省得每个脚本单独配。入口在https://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_contentcyclegan_tf。3. 可复制的 CycleGAN 环境依赖与训练配置这一节给全可复制的配置。先建虚拟环境Python 3.9–3.11 都行我实测 3.10 最稳。python -m venv cyclegan_env source cyclegan_env/bin/activate # Windows: cyclegan_env\Scripts\activate pip install --upgrade pip pip install tensorflow2.15.0 tensorflow-datasets4.9.4 tensorflow-examples0.0.1 pip install numpy1.24.3 matplotlib3.7.2 scipy1.11.1 pip install githttps://github.com/tensorflow/examples.git#eggtensorflow-examplestensorflow-examples这个包有时候 pip 源里版本对不上直接用 git 装最保险。装完验证一下import tensorflow as tf from tensorflow_examples.models.pix2pix import pix2pix print(tf.__version__) print(pix2pix.unet_generator)能打印出function unet_generator at ...就说明生成器模块可用。接下来是项目配置文件。我用一个config.yaml把路径、超参、凭据引用都收在一起避免散落在代码里# config.yaml project: name: cyclegan_apple2orange seed: 42 output_dir: ./outputs data: dataset_name: apple2orange data_root: ./data img_height: 256 img_width: 256 batch_size: 1 buffer_size: 1000 train: epochs: 100 lambda_cycle: 10.0 lambda_identity: 5.0 lr: 0.0002 beta_1: 0.5 checkpoint_dir: ./checkpoints save_every: 5 sample_every: 1 api: base_url: https://taotoken.net/api api_key_env: TAOTOKEN_API_KEY model_id: your-model-id注意api段里 Key 不写明文只写环境变量名脚本运行时去读。model_id按你实际要调的服务填。然后是数据加载和预处理脚本data_loader.pyimport tensorflow as tf from config import load_config cfg load_config(config.yaml) AUTOTUNE tf.data.AUTOTUNE IMG_H cfg[data][img_height] IMG_W cfg[data][img_width] BATCH cfg[data][batch_size] BUFFER cfg[data][buffer_size] def load_image(path): image tf.io.read_file(path) image tf.image.decode_jpeg(image, channels3) return tf.cast(image, tf.float32) def normalize(image): return (image / 127.5) - 1.0 def random_jitter(image): image tf.image.resize(image, [286, 286], methodtf.image.ResizeMethod.NEAREST_NEIGHBOR) image tf.image.random_crop(image, size[IMG_H, IMG_W, 3]) image tf.image.random_flip_left_right(image) return image def preprocess_train(path): image load_image(path) image random_jitter(image) return normalize(image) def preprocess_test(path): image load_image(path) return normalize(image) def build_datasets(data_root, dataset_name): train_a tf.data.Dataset.list_files( f{data_root}/{dataset_name}/trainA/*.jpg, seedcfg[project][seed]) train_b tf.data.Dataset.list_files( f{data_root}/{dataset_name}/trainB/*.jpg, seedcfg[project][seed]) test_a tf.data.Dataset.list_files( f{data_root}/{dataset_name}/testA/*.jpg, seedcfg[project][seed]) test_b tf.data.Dataset.list_files( f{data_root}/{dataset_name}/testB/*.jpg, seedcfg[project][seed]) train_a train_a.map(preprocess_train, num_parallel_callsAUTOTUNE) \ .shuffle(BUFFER, seedcfg[project][seed]) \ .batch(BATCH, drop_remainderTrue).prefetch(AUTOTUNE) train_b train_b.map(preprocess_train, num_parallel_callsAUTOTUNE) \ .shuffle(BUFFER, seedcfg[project][seed]) \ .batch(BATCH, drop_remainderTrue).prefetch(AUTOTUNE) test_a test_a.map(preprocess_test, num_parallel_callsAUTOTUNE) \ .cache().batch(BATCH, drop_remainderTrue).prefetch(AUTOTUNE) test_b test_b.map(preprocess_test, num_parallel_callsAUTOTUNE) \ .cache().batch(BATCH, drop_remainderTrue).prefetch(AUTOTUNE) return train_a, train_b, test_a, test_b这里的关键点是seed固定保证每次跑的数据顺序一致方便复现。drop_remainderTrue避免最后一个不满 batch 的样本干扰 BatchNorm 统计。模型定义脚本models.pyimport tensorflow as tf from tensorflow_examples.models.pix2pix import pix2pix OUTPUT_CHANNELS 3 def build_generators(): gen_g pix2pix.unet_generator(OUTPUT_CHANNELS, norm_typeinstancenorm) gen_f pix2pix.unet_generator(OUTPUT_CHANNELS, norm_typeinstancenorm) return gen_g, gen_f def build_discriminators(): disc_x pix2pix.discriminator(norm_typeinstancenorm, targetFalse) disc_y pix2pix.discriminator(norm_typeinstancenorm, targetFalse) return disc_x, disc_y def build_optimizers(lr2e-4, beta_10.5): return ( tf.keras.optimizers.Adam(lr, beta_1beta_1), tf.keras.optimizers.Adam(lr, beta_1beta_1), tf.keras.optimizers.Adam(lr, beta_1beta_1), tf.keras.optimizers.Adam(lr, beta_1beta_1), )损失函数和训练步train_step.pyimport tensorflow as tf LAMBDA_CYCLE 10.0 LAMBDA_IDENTITY 5.0 loss_obj tf.keras.losses.BinaryCrossentropy(from_logitsTrue) def discriminator_loss(real, generated): real_loss loss_obj(tf.ones_like(real), real) gen_loss loss_obj(tf.zeros_like(generated), generated) return (real_loss gen_loss) * 0.5 def generator_loss(generated): return loss_obj(tf.ones_like(generated), generated) def calc_cycle_loss(real_image, cycled_image): return LAMBDA_CYCLE * tf.reduce_mean(tf.abs(real_image - cycled_image)) def identity_loss(real_image, same_image): return LAMBDA_IDENTITY * 0.5 * tf.reduce_mean(tf.abs(real_image - same_image)) tf.function def train_step(real_x, real_y, gen_g, gen_f, disc_x, disc_y, opt_g, opt_f, opt_dx, opt_dy): with tf.GradientTape(persistentTrue) as tape: fake_y gen_g(real_x, trainingTrue) cycled_x gen_f(fake_y, trainingTrue) fake_x gen_f(real_y, trainingTrue) cycled_y gen_g(fake_x, trainingTrue) same_x gen_f(real_x, trainingTrue) same_y gen_g(real_y, trainingTrue) disc_real_x disc_x(real_x, trainingTrue) disc_real_y disc_y(real_y, trainingTrue) disc_fake_x disc_x(fake_x, trainingTrue) disc_fake_y disc_y(fake_y, trainingTrue) gen_g_loss generator_loss(disc_fake_y) gen_f_loss generator_loss(disc_fake_x) total_cycle calc_cycle_loss(real_x, cycled_x) \ calc_cycle_loss(real_y, cycled_y) total_gen_g gen_g_loss total_cycle identity_loss(real_y, same_y) total_gen_f gen_f_loss total_cycle identity_loss(real_x, same_x) disc_x_loss discriminator_loss(disc_real_x, disc_fake_x) disc_y_loss discriminator_loss(disc_real_y, disc_fake_y) grads_g tape.gradient(total_gen_g, gen_g.trainable_variables) grads_f tape.gradient(total_gen_f, gen_f.trainable_variables) grads_dx tape.gradient(disc_x_loss, disc_x.trainable_variables) grads_dy tape.gradient(disc_y_loss, disc_y.trainable_variables) opt_g.apply_gradients(zip(grads_g, gen_g.trainable_variables)) opt_f.apply_gradients(zip(grads_f, gen_f.trainable_variables)) opt_dx.apply_gradients(zip(grads_dx, disc_x.trainable_variables)) opt_dy.apply_gradients(zip(grads_dy, disc_y.trainable_variables)) return total_gen_g, total_gen_f, disc_x_loss, disc_y_losspersistentTrue是必须的因为同一个 tape 要对四组变量分别求梯度。漏了这个参数第二次tape.gradient就会报 “GradientTape.gradient can only be called once”。主训练脚本train.py把上面串起来加上 checkpoint 和固定种子import os, time, random import numpy as np import tensorflow as tf from data_loader import build_datasets from models import build_generators, build_discriminators, build_optimizers from train_step import train_step SEED 42 random.seed(SEED) np.random.seed(SEED) tf.random.set_seed(SEED) def main(): train_a, train_b, test_a, test_b build_datasets(./data, apple2orange) gen_g, gen_f build_generators() disc_x, disc_y build_discriminators() opt_g, opt_f, opt_dx, opt_dy build_optimizers() ckpt tf.train.Checkpoint( gen_ggen_g, gen_fgen_f, disc_xdisc_x, disc_ydisc_y, opt_gopt_g, opt_fopt_f, opt_dxopt_dx, opt_dyopt_dy) ckpt_manager tf.train.CheckpointManager( ckpt, ./checkpoints, max_to_keep5) if ckpt_manager.latest_checkpoint: ckpt.restore(ckpt_manager.latest_checkpoint) print(fRestored from {ckpt_manager.latest_checkpoint}) EPOCHS 100 for epoch in range(EPOCHS): start time.time() n 0 for real_x, real_y in tf.data.Dataset.zip((train_a, train_b)): g_loss, f_loss, dx_loss, dy_loss train_step( real_x, real_y, gen_g, gen_f, disc_x, disc_y, opt_g, opt_f, opt_dx, opt_dy) if n % 50 0: print(fEpoch {epoch1} step {n} | fG:{g_loss:.3f} F:{f_loss:.3f} fDx:{dx_loss:.3f} Dy:{dy_loss:.3f}) n 1 if (epoch 1) % 5 0: path ckpt_manager.save() print(fSaved checkpoint: {path}) print(fEpoch {epoch1} done in {time.time()-start:.1f}s) if __name__ __main__: main()数据集下载用tensorflow_datasets或者直接下官方 zip。apple2orange官方包大概 300MB解压后目录是apple2orange/trainA、trainB、testA、testB和上面代码的路径约定一致。4. 验证请求与成功结果固定种子下的 FID 与视觉样例训练跑起来之后怎么判断模型真的学到了迁移而不是在输出噪声光看 loss 不够CycleGAN 的 loss 曲线经常看起来很平稳但生成质量很差。我用两个手段交叉验证固定种子的视觉样例 FID 指标。视觉样例脚本sample.pyimport matplotlib.pyplot as plt import tensorflow as tf from data_loader import build_datasets from models import build_generators def generate_and_save(model, test_input, save_path): prediction model(test_input, trainingFalse) plt.figure(figsize(10, 5)) display [test_input[0], prediction[0]] titles [Input, Predicted] for i in range(2): plt.subplot(1, 2, i 1) plt.title(titles[i]) plt.imshow(display[i] * 0.5 0.5) plt.axis(off) plt.savefig(save_path, dpi150, bbox_inchestight) plt.close() def main(): _, _, test_a, test_b build_datasets(./data, apple2orange) gen_g, gen_f build_generators() ckpt tf.train.Checkpoint(gen_ggen_g, gen_fgen_f) ckpt.restore(tf.train.latest_checkpoint(./checkpoints)).expect_partial() for i, inp in enumerate(test_a.take(5)): generate_and_save(gen_g, inp, f./outputs/apple_to_orange_{i}.png) for i, inp in enumerate(test_b.take(5)): generate_and_save(gen_f, inp, f./outputs/orange_to_apple_{i}.png) print(Samples saved to ./outputs/) if __name__ __main__: main()跑完打开./outputs/里的图苹果应该变成橘子的暖色调形状结构保留反过来橘子变苹果的冷绿色调。如果输出是全灰或者严重棋盘伪影说明训练有问题往下看第 5 节。FID 计算用scipy和预训练的 InceptionV3 特征。这里给一个轻量实现import numpy as np import tensorflow as tf from scipy.linalg import sqrtm def get_inception_features(images, batch_size32): inception tf.keras.applications.InceptionV3( include_topFalse, poolingavg, input_shape(299, 299, 3)) feats [] for i in range(0, len(images), batch_size): batch images[i:i batch_size] batch tf.image.resize(batch, (299, 299)) batch tf.keras.applications.inception_v3.preprocess_input(batch) feats.append(inception(batch, trainingFalse).numpy()) return np.concatenate(feats, axis0) def calculate_fid(real_images, fake_images): real_feat get_inception_features(real_images) fake_feat get_inception_features(fake_images) mu_r, sigma_r real_feat.mean(0), np.cov(real_feat, rowvarFalse) mu_f, sigma_f fake_feat.mean(0), np.cov(fake_feat, rowvarFalse) diff mu_r - mu_f covmean sqrtm(sigma_r sigma_f) if np.iscomplexobj(covmean): covmean covmean.real fid diff diff np.trace(sigma_r sigma_f - 2 * covmean) return float(fid)实测下来苹果↔橘子数据集上训练 60 epoch 后 FID 大概能降到 80–110 区间具体数值和随机种子、数据划分有关不要拿这个当绝对标准。关键是看趋势从 epoch 20 的 200 降到 epoch 60 的 100 左右说明模型在收敛。如果 FID 一直不降甚至上升多半是判别器太强或学习率不对。固定种子的意义在这里体现每次跑sample.py拿到的输入图是同一批生成的对比图可以直接叠着看判断是模型进步了还是数据换了。没有固定种子你根本分不清变化来自哪里。5. 本篇常见报错排查401、local proxy failed、reading choices、OAuth这一节按真实报错来。CycleGAN 训练本身不涉及网络请求但你的实验管理脚本、FID 评估辅助、日志摘要生成这些环节会调 TaoToken 通道报错集中在这几类。401 Unauthorized。最常见的原因是 Key 没读到或读错。检查顺序先确认环境变量真的导出了echo $TAOTOKEN_API_KEYWindows 用echo $env:TAOTOKEN_API_KEY看有没有值再确认脚本里读的是同一个变量名别一个写TAOTOKEN_API_KEY另一个写TAOTOKEN_KEY最后确认 Key 没有多余空格或换行复制时容易带上。如果用的是.env文件确认load_dotenv()在读取配置之前调用。local proxy failed。这个报错通常出现在你的运行环境里配置了本地网络设置导致请求发不出去。TaoToken 的 API 通道是标准 HTTPS 直连不需要任何额外网络配置。检查HTTP_PROXY、HTTPS_PROXY、ALL_PROXY这几个环境变量如果被设成了本地地址清掉unset HTTP_PROXY HTTPS_PROXY ALL_PROXYWindows 下在“环境变量”设置里删掉对应的用户变量。清完之后重新跑请求。reading choices 相关报错。这个一般出现在解析 API 返回的 JSON 时字段路径写错了。TaoToken 的对话接口返回结构里内容在choices[0].message.content如果你按别的路径取就会报 KeyError 或 reading 失败。建议先打印完整响应import os, requests, json resp requests.post( f{os.environ[TAOTOKEN_BASE_URL]}/v1/chat/completions, headers{Authorization: fBearer {os.environ[TAOTOKEN_API_KEY]}, Content-Type: application/json}, json{model: your-model-id, messages: [{role: user, content: test}]}, timeout30) print(resp.status_code) print(json.dumps(resp.json(), ensure_asciiFalse, indent2))看清楚结构再写解析代码。OAuth 相关报错。如果你用的是需要 OAuth 流程的客户端比如某些 IDE 插件或 CLI 工具报 OAuth 失败通常是回调地址或 token 过期。TaoToken 的 API Key 方式是 Bearer Token不涉及 OAuth 跳转。如果你在某个工具里看到 OAuth 报错检查是不是工具本身配置了别的认证方式把它切回 API Key 模式Base URL 填https://taotoken.net/apiKey 填你的实际 KeyModel ID 填对应模型。还有一个容易忽略的checkpoint 恢复后 loss 跳变。这不是 API 报错但很常见。原因是tf.train.Checkpoint里如果只存了模型没存优化器恢复后 Adam 的动量状态归零前几个 step 的更新幅度会异常。上面的train.py里我把四个优化器都放进了 Checkpoint就是为了避免这个问题。如果你自己写的时候漏了恢复后 loss 突然飙高是正常现象跑几十个 step 会稳回来但最好一开始就存全。生成器输出全灰。训练早期正常如果 20 epoch 后还是灰的检查normalize是不是把图像归到了 [-1,1]以及可视化时有没有做* 0.5 0.5反归一化。另一个可能是LAMBDA_IDENTITY设太大生成器倾向于恒等映射不做事把它从 5.0 降到 2.0 试试。6. 把 CycleGAN 接入你的实验流水线训练脚本跑通只是第一步。真正让这套东西有价值的是把它接进你的实验流水线数据版本管理、训练日志、模型评估、结果归档。TaoToken 在这里的价值是让这些环节的模型调用有统一凭据不用每个脚本单独配。具体做法在项目根目录放一个api_client.py封装所有对外调用import os import requests class TaoTokenClient: def __init__(self): self.base_url os.environ.get(TAOTOKEN_BASE_URL, https://taotoken.net/api) self.api_key os.environ[TAOTOKEN_API_KEY] self.model_id os.environ.get(TAOTOKEN_MODEL_ID, your-model-id) def chat(self, prompt, timeout60): resp requests.post( f{self.base_url}/v1/chat/completions, headers{Authorization: fBearer {self.api_key}, Content-Type: application/json}, json{model: self.model_id, messages: [{role: user, content: prompt}]}, timeouttimeout) resp.raise_for_status() return resp.json()[choices][0][message][content]然后在训练脚本的 epoch 回调里调它生成训练摘要或者用 coding plan 通道跑实验脚本的辅助生成。这样你的 CycleGAN 项目就有了一个统一的模型能力入口而不是散落各处的硬编码。如果你要跑多个数据集对比苹果↔橘子、马↔斑马、夏↔冬建议把config.yaml里的dataset_name参数化用命令行覆盖python train.py --dataset horse2zebra --epochs 150配合argparse读参数一套代码跑所有数据集。checkpoint 目录按数据集名分开存避免互相覆盖。最后给一个实用技巧CycleGAN 训练到后期判别器和生成器的平衡很微妙。如果发现生成图像开始出现明显伪影把判别器的学习率降到生成器的一半或者给判别器加一点输入噪声。这个调整不需要改网络结构在build_optimizers里给opt_dx、opt_dy传不同的lr就行。实测在马↔斑马这种纹理变化大的数据集上这个调整能让 FID 再降 10–15 个点。整套流程跑下来从环境配置到出第一张可看的迁移图单卡大概 6–8 小时100 epoch。如果你只想快速验证流程通不通把EPOCHS改成 5跑完看./outputs/里有没有图出来有图且不是全灰就说明整条链路没问题剩下的就是等它慢慢收敛。
返回列表