用 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 Network,GAN)是一类深度学习模型,由 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, nc=4, nz=100, ngf=64): super(Generator, self).__init__() self.network = nn.Sequential( nn.ConvTranspose2d(nz, ngf * 4, 3, 1, 0, bias=False), nn.BatchNorm2d(ngf * 4), nn.ReLU(True), nn.ConvTranspose2d(ngf * 4, ngf * 2, 3, 2, 1, bias=False), nn.BatchNorm2d(ngf * 2), nn.ReLU(True), nn.ConvTranspose2d(ngf * 2, ngf, 4, 2, 0, bias=False), nn.BatchNorm2d(ngf), nn.ReLU(True), nn.ConvTranspose2d(ngf, nc, 4, 2, 1, bias=False), nn.Tanh(), ) def forward(self, input): output = self.network(input) return output对三个关键超参数的理解有助于你改造成其他生成器:
nc:输出图像的通道数(CryptoPunks 像素画为 RGBA,故取 4;普通 RGB 图像为 3);nz:输入噪声向量的长度(本模型为 100,是 DCGAN 常用取值);ngf:生成器特征图通道数的基准倍数,决定网络宽度。
结构上,网络由四层ConvTranspose2d(转置卷积,负责把低分辨率特征逐步放大)配合BatchNorm2d+ReLU组成,最后一层以Tanh将像素值约束到[-1, 1],这与torchvision的save_image配合时可使用normalize=True还原显示。可见层数越深、通道扩展规律越清晰,模型就越接近标准 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_location=torch.device('cpu'))) # 如果有可用的GPU,请使用'cuda'要点:
hf_hub_download('nateraw/cryptopunks-gan', 'generator.pth')会从 Hub 下载指定文件并返回本地缓存路径,之后无需手动管理下载目录;map_location=torch.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", normalize=True) return 'punks.png'设计细节值得展开:
seed参数:通过torch.manual_seed(seed)固定随机数生成。由于torch.randn的随机性依赖种子,传入相同seed会得到相同的噪声张量,从而可以稳定复现同一批 punk 图像;- 输入张量维度:模型要求单次推理输入为
100x1x1,批量推理为(BatchSize)x100x1x1。本例每次生成 4 个 punk,因此张量为(4, 100, 1, 1); - 输出:
save_image会把张量网格化为一张 PNG,normalize=True会把网络输出的[-1, 1]区间归一化后再写盘。函数最终返回图片文件路径字符串'punks.png',Gradio 的 Image 输出组件可直接渲染该路径对应的图片。
第三步:创建 Gradio 接口——一个函数调用定义整个应用
到这一步,直接运行predict(<某个数字>)已经能在文件系统./punks.png中找到新生成的 punk 图。但要做出真正可交互的演示,还需要一个界面。目标拆解为三点:
- 一个滑块输入,让用户自由选择
seed值; - 一个图像输出组件,用于展示生成的 punk 图;
- 由
predict()承接「取种子 → 生成图像」的完整链路。
借助gr.Interface,一次函数调用即可描述以上全部内容:
import gradio as gr gr.Interface( predict, inputs=[ gr.Slider(0, 1000, label='Seed', default=42), ], outputs="image", ).launch()inputs列表中的gr.Slider(0, 1000, ...)定义取值范围为[0, 1000]、默认值 42 的整数种子滑块,其显示名称由label指定;outputs="image"表示输出使用 Image 组件(Gradio 允许用简短字符串快捷指定组件,与gr.Image()等价),它会自动处理predict返回的文件路径;- 由于
predict(seed)恰好只有一个入参、一个返回值,与inputs/outputs一一对应,Gradio 会自动完成参数绑定。
launch()启动后应用即在本地运行,打开浏览器即可拖动滑块看到不同种子对应的 punk 图。
第四步:加入「数量」滑块——多输入与函数签名的联动
每次固定生成 4 个 punk 是个不错的起点,但若想自由控制每次生成数量,只需向inputs列表追加一项输入:
gr.Interface( predict, inputs=[ gr.Slider(0, 1000, label='Seed', default=42), gr.Slider(4, 64, label='Number of Punks', step=1, default=10), # 添加另一个滑块! ], outputs="image", ).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", normalize=True) return 'punks.png'注意事项:
- 第二个滑块用
step=1限定整数取值(数量没有小数意义),范围[4, 64]与模型可一次性生成的上限相匹配; - 多输入场景下,
predict的形参顺序必须与inputs列表中组件的排列顺序严格一致——这是 Gradio 事件触发的隐式契约,顺序错位会造成参数张冠李戴; - 重启界面后即可看到第二个滑块,实时控制每次生成的 punk 数量。
第五步:打磨体验——examples、标题与描述性内容
应用功能已可用,再加几个小功能就能让最终效果更出彩。
添加一键示例
examples参数允许预设一组可点击即用的输入组合,用户无需手动拖滑块即可快速体验:
gr.Interface( # ... # 将所有内容保持不变,然后添加 examples=[[123, 15], [42, 29], [456, 8], [1337, 35]], ).launch(cache_examples=True) # cache_examples是可选的examples接受一个列表的列表,每个子列表的条目顺序与inputs声明的顺序一致,即本例中的[seed, num_punks];- 界面中每个示例会以缩略卡片形式展示,点击即自动填充两个滑块并触发预测;
cache_examples=True会启动时预先缓存示例运行结果(此时需保证函数可离线执行),可显著加速用户点击示例后的响应;若不设置则每次实时计算。
添加标题、描述与署名内容
可以为gr.Interface添加title、description与article,三者均接受字符串:
title:显示在界面顶部,同时作为浏览器页面标题;description:放置在标题正下方,可接受文本、Markdown 或 HTML;article:放置在界面下方,同样可接受文本、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, nc=4, nz=100, ngf=64): super(Generator, self).__init__() self.network = nn.Sequential( nn.ConvTranspose2d(nz, ngf * 4, 3, 1, 0, bias=False), nn.BatchNorm2d(ngf * 4), nn.ReLU(True), nn.ConvTranspose2d(ngf * 4, ngf * 2, 3, 2, 1, bias=False), nn.BatchNorm2d(ngf * 2), nn.ReLU(True), nn.ConvTranspose2d(ngf * 2, ngf, 4, 2, 0, bias=False), nn.BatchNorm2d(ngf), nn.ReLU(True), nn.ConvTranspose2d(ngf, nc, 4, 2, 1, bias=False), 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_location=torch.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", normalize=True) return 'punks.png' gr.Interface( predict, inputs=[ gr.Slider(0, 1000, label='Seed', default=42), gr.Slider(4, 64, label='Number of Punks', step=1, default=10), ], outputs="image", examples=[[123, 15], [42, 29], [456, 8], [1337, 35]], ).launch(cache_examples=True)代码组织上可归纳为清晰的三段式模板,适合复用到其他生成模型:
- 模型区:定义网络结构 + 从 Hub 加载权重;
- 预测区:把「随机种子 → 输出文件路径」封装为与界面输入一一对应的纯函数;
- 界面区:用
gr.Interface声明输入、输出、示例与描述性内容。
延伸:把模板迁移到你的模型
这个 Demo 的意义远超「生成 CryptoPunks」本身。剥离掉领域细节后,它沉淀出一个通用的**「预训练生成式模型快速 Demo 化」模板**:
- 换掉
Generator结构、换一个 Hub 模型仓库(如各类 GAN、VAE、扩散模型的生成器权重),保留「随机噪声/潜变量 → 图片」的范式即可复用全部 Gradio 代码; - 若你的模型输入不是整数种子而是文本或图片,把
predict的入参与inputs列表替换为对应的gr.Textbox、gr.Image等组件即可,函数签名联动逻辑不变; - 若模型输出为多张图网格,可借助 Image 输出网格;若涉及批量逐张展示,则可参考仓库中基于
gr.Blocks+gr.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),仅供参考