资讯动态

用 Gradio 为预训练 GAN 构建 CryptoPunks 生成器:从模型加载到可交互 Web 应用全流程

发布时间:2026/9/10 19:37:31 来源:尧图企业网站定制
用 Gradio 为预训练 GAN 构建 CryptoPunks 生成器从模型加载到可交互 Web 应用全流程【免费下载链接】gradioBuild and share delightful machine learning apps, all in Python. Star to support our work!项目地址: https://gitcode.com/GitHub_Trending/gr/gradio本文基于 Gradio 官方中文教程guides/cn/07_other-tutorials/create-your-own-friends-with-a-gan.md讲解如何把一个「只能输出图片文件的 PyTorch 生成器模型」包装成一个任何人都能通过浏览器使用的交互式生成器。你将掌握GAN 生成器模型如何定义与加载权重、predict函数如何连接模型与界面、gr.Interface如何用滑块输入与图像输出快速搭建演示以及如何通过examples和描述性参数把粗糙原型打磨成可分享的成品。整个思路同样适用于将任何预训练生成式模型图像、音频、文本快速 Demo 化。背景什么是 GAN为什么只需「生成器」生成对抗网络Generative Adversarial NetworkGAN是一类深度学习模型由 Goodfellow 等人在 2014 年提出。它由两个相互竞争的神经网络组成生成器Generator负责从随机噪声中生成图像鉴别器Discriminator接收生成器产出的图片与训练集中的真实图片并判断哪张是伪造的。生成器不断学习如何制造更难被识别的图像而鉴别器每识破一张假图就相当于为生成器提高了门槛。随着这种「对抗」式训练持续推进生成图像的质量会逐步提升到接近以假乱真的水平。一个关键推论是在推理生成新图像阶段只需要生成器模型鉴别器仅在训练阶段发挥作用。这正是本教程的实操基础——我们下载一个预训练生成器跳过训练直接做生成。为了直观理解可以参考仓库中的 demo/fake_gan/run.py它用gr.Blocks 随机选图模拟了一个「假 GAN」界面展示了真实 GAN demo 在 Gradio 中的典型交互形态按钮触发、Gallery 展示生成结果。说明本篇为教程向内容涉及行业背景仅作科普铺垫实际动手所需的环境依赖、代码与组件行为均以当前仓库为准。前置条件安装依赖开始前需要确保环境满足Python 环境已安装gradio包安装方式参见中文快速入门指南 guides/cn/01_getting-started/01_quickstart.md由于要加载并运行 PyTorch 预训练模型还需额外安装torch与torchvision从 Hugging Face Hub 拉取权重依赖huggingface_hub新版torch/huggingface_hub环境一般已内置。gradio的Interface、Slider、Image等 API 由仓库内 gradio/interface.py、gradio/components/slider.py、gradio/components/image.py 实现读者可随时对照源码确认参数行为。第一步定义生成器模型并加载预训练权重本教程使用的生成器是一个典型的 DCGAN 风格反卷积网络将 100 维随机噪声向量上采样为一张小尺寸图像。代码来自公开的 CryptoPunks GAN 训练仓库模型权重发布在 Hugging Face Hub 的nateraw/cryptopunks-gan仓库文件名为generator.pth。模型结构from torch import nn class Generator(nn.Module): # 有关 nc、nz 和 ngf 的解释请参见 DCGAN 官方教程的 Inputs 小节 def __init__(self, nc4, nz100, ngf64): super(Generator, self).__init__() self.network nn.Sequential( nn.ConvTranspose2d(nz, ngf * 4, 3, 1, 0, biasFalse), nn.BatchNorm2d(ngf * 4), nn.ReLU(True), nn.ConvTranspose2d(ngf * 4, ngf * 2, 3, 2, 1, biasFalse), nn.BatchNorm2d(ngf * 2), nn.ReLU(True), nn.ConvTranspose2d(ngf * 2, ngf, 4, 2, 0, biasFalse), nn.BatchNorm2d(ngf), nn.ReLU(True), nn.ConvTranspose2d(ngf, nc, 4, 2, 1, biasFalse), nn.Tanh(), ) def forward(self, input): output self.network(input) return output对三个关键超参数的理解有助于你改造成其他生成器nc输出图像的通道数CryptoPunks 像素画为 RGBA故取 4普通 RGB 图像为 3nz输入噪声向量的长度本模型为 100是 DCGAN 常用取值ngf生成器特征图通道数的基准倍数决定网络宽度。结构上网络由四层ConvTranspose2d转置卷积负责把低分辨率特征逐步放大配合BatchNorm2dReLU组成最后一层以Tanh将像素值约束到[-1, 1]这与torchvision的save_image配合时可使用normalizeTrue还原显示。可见层数越深、通道扩展规律越清晰模型就越接近标准 DCGAN——这是推理端可直接复用、无需训练的标志。加载预训练权重from huggingface_hub import hf_hub_download import torch model Generator() weights_path hf_hub_download(nateraw/cryptopunks-gan, generator.pth) model.load_state_dict(torch.load(weights_path, map_locationtorch.device(cpu))) # 如果有可用的GPU请使用cuda要点hf_hub_download(nateraw/cryptopunks-gan, generator.pth)会从 Hub 下载指定文件并返回本地缓存路径之后无需手动管理下载目录map_locationtorch.device(cpu)强制将权重载入 CPU若机器有 GPU可改为cuda以加速生成若只生成本文的model在下载到本地前会缓存代码可重复运行而无需重复下载。第二步定义predict函数——Gradio 应用的心脏predict函数是让 Gradio 运转起来的关键用户在界面中做出的任何输入都会作为参数传入predict其返回值再交由 Gradio 输出组件渲染。对 GAN 来说惯例是把随机噪声作为模型输入因此我们生成一个随机数张量送入模型再用torchvision的save_image把输出保存为 PNG 文件并返回文件名from torchvision.utils import save_image def predict(seed): num_punks 4 torch.manual_seed(seed) z torch.randn(num_punks, 100, 1, 1) punks model(z) save_image(punks, punks.png, normalizeTrue) return punks.png设计细节值得展开seed参数通过torch.manual_seed(seed)固定随机数生成。由于torch.randn的随机性依赖种子传入相同seed会得到相同的噪声张量从而可以稳定复现同一批 punk 图像输入张量维度模型要求单次推理输入为100x1x1批量推理为(BatchSize)x100x1x1。本例每次生成 4 个 punk因此张量为(4, 100, 1, 1)输出save_image会把张量网格化为一张 PNGnormalizeTrue会把网络输出的[-1, 1]区间归一化后再写盘。函数最终返回图片文件路径字符串punks.pngGradio 的 Image 输出组件可直接渲染该路径对应的图片。第三步创建 Gradio 接口——一个函数调用定义整个应用到这一步直接运行predict(某个数字)已经能在文件系统./punks.png中找到新生成的 punk 图。但要做出真正可交互的演示还需要一个界面。目标拆解为三点一个滑块输入让用户自由选择seed值一个图像输出组件用于展示生成的 punk 图由predict()承接「取种子 → 生成图像」的完整链路。借助gr.Interface一次函数调用即可描述以上全部内容import gradio as gr gr.Interface( predict, inputs[ gr.Slider(0, 1000, labelSeed, default42), ], outputsimage, ).launch()inputs列表中的gr.Slider(0, 1000, ...)定义取值范围为[0, 1000]、默认值 42 的整数种子滑块其显示名称由label指定outputsimage表示输出使用 Image 组件Gradio 允许用简短字符串快捷指定组件与gr.Image()等价它会自动处理predict返回的文件路径由于predict(seed)恰好只有一个入参、一个返回值与inputs/outputs一一对应Gradio 会自动完成参数绑定。launch()启动后应用即在本地运行打开浏览器即可拖动滑块看到不同种子对应的 punk 图。第四步加入「数量」滑块——多输入与函数签名的联动每次固定生成 4 个 punk 是个不错的起点但若想自由控制每次生成数量只需向inputs列表追加一项输入gr.Interface( predict, inputs[ gr.Slider(0, 1000, labelSeed, default42), gr.Slider(4, 64, labelNumber of Punks, step1, default10), # 添加另一个滑块! ], outputsimage, ).launch()新增输入会按照声明顺序自动传递给predict()因此函数签名也必须同步增加一个参数def predict(seed, num_punks): torch.manual_seed(seed) z torch.randn(num_punks, 100, 1, 1) punks model(z) save_image(punks, punks.png, normalizeTrue) return punks.png注意事项第二个滑块用step1限定整数取值数量没有小数意义范围[4, 64]与模型可一次性生成的上限相匹配多输入场景下predict的形参顺序必须与inputs列表中组件的排列顺序严格一致——这是 Gradio 事件触发的隐式契约顺序错位会造成参数张冠李戴重启界面后即可看到第二个滑块实时控制每次生成的 punk 数量。第五步打磨体验——examples、标题与描述性内容应用功能已可用再加几个小功能就能让最终效果更出彩。添加一键示例examples参数允许预设一组可点击即用的输入组合用户无需手动拖滑块即可快速体验gr.Interface( # ... # 将所有内容保持不变然后添加 examples[[123, 15], [42, 29], [456, 8], [1337, 35]], ).launch(cache_examplesTrue) # cache_examples是可选的examples接受一个列表的列表每个子列表的条目顺序与inputs声明的顺序一致即本例中的[seed, num_punks]界面中每个示例会以缩略卡片形式展示点击即自动填充两个滑块并触发预测cache_examplesTrue会启动时预先缓存示例运行结果此时需保证函数可离线执行可显著加速用户点击示例后的响应若不设置则每次实时计算。添加标题、描述与署名内容可以为gr.Interface添加title、description与article三者均接受字符串title显示在界面顶部同时作为浏览器页面标题description放置在标题正下方可接受文本、Markdown 或 HTMLarticle放置在界面下方同样可接受文本、Markdown 或 HTML。article支持 HTML 的细节以及 Blocks 中如何使用gr.Markdown/gr.HTML内联描述性内容详见中文指南「描述性内容」章节 guides/cn/01_getting-started/02_key-features.md。完整代码参考以下是教程全部代码的汇总运行后浏览器打开本地地址即可交互import torch from torch import nn from huggingface_hub import hf_hub_download from torchvision.utils import save_image import gradio as gr class Generator(nn.Module): # 关于 nc、nz 和 ngf 的解释请参见 DCGAN 官方教程的 Inputs 小节 def __init__(self, nc4, nz100, ngf64): super(Generator, self).__init__() self.network nn.Sequential( nn.ConvTranspose2d(nz, ngf * 4, 3, 1, 0, biasFalse), nn.BatchNorm2d(ngf * 4), nn.ReLU(True), nn.ConvTranspose2d(ngf * 4, ngf * 2, 3, 2, 1, biasFalse), nn.BatchNorm2d(ngf * 2), nn.ReLU(True), nn.ConvTranspose2d(ngf * 2, ngf, 4, 2, 0, biasFalse), nn.BatchNorm2d(ngf), nn.ReLU(True), nn.ConvTranspose2d(ngf, nc, 4, 2, 1, biasFalse), nn.Tanh(), ) def forward(self, input): output self.network(input) return output model Generator() weights_path hf_hub_download(nateraw/cryptopunks-gan, generator.pth) model.load_state_dict(torch.load(weights_path, map_locationtorch.device(cpu))) # 如果您有可用的GPU使用cuda def predict(seed, num_punks): torch.manual_seed(seed) z torch.randn(num_punks, 100, 1, 1) punks model(z) save_image(punks, punks.png, normalizeTrue) return punks.png gr.Interface( predict, inputs[ gr.Slider(0, 1000, labelSeed, default42), gr.Slider(4, 64, labelNumber of Punks, step1, default10), ], outputsimage, examples[[123, 15], [42, 29], [456, 8], [1337, 35]], ).launch(cache_examplesTrue)代码组织上可归纳为清晰的三段式模板适合复用到其他生成模型模型区定义网络结构 从 Hub 加载权重预测区把「随机种子 → 输出文件路径」封装为与界面输入一一对应的纯函数界面区用gr.Interface声明输入、输出、示例与描述性内容。延伸把模板迁移到你的模型这个 Demo 的意义远超「生成 CryptoPunks」本身。剥离掉领域细节后它沉淀出一个通用的**「预训练生成式模型快速 Demo 化」模板**换掉Generator结构、换一个 Hub 模型仓库如各类 GAN、VAE、扩散模型的生成器权重保留「随机噪声/潜变量 → 图片」的范式即可复用全部 Gradio 代码若你的模型输入不是整数种子而是文本或图片把predict的入参与inputs列表替换为对应的gr.Textbox、gr.Image等组件即可函数签名联动逻辑不变若模型输出为多张图网格可借助 Image 输出网格若涉及批量逐张展示则可参考仓库中基于gr.Blocksgr.Gallery的交互写法见 demo/fake_gan/run.py构造更灵活的布局。本文对应的英文原版教程位于 guides/11_other-tutorials/create-your-own-friends-with-a-gan.md供对照阅读。恭喜至此你已经完成了一个具备滑块输入、图像输出、示例一键复现的 GAN 生成器应用完全可以在此基础上持续挖掘 Hub 上更多生成式模型打造更多演示项目。【免费下载链接】gradioBuild and share delightful machine learning apps, all in Python. Star to support our work!项目地址: https://gitcode.com/GitHub_Trending/gr/gradio创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价