如何使用 annotated_deep_learning_paper_implementations 的 Stable Diffusion 脚本完成文本到图像生成
【免费下载链接】annotated_deep_learning_paper_implementations🧑🏫 60+ Implementations/tutorials of deep learning papers with side-by-side notes 📝; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), 🎮 reinforcement learning (ppo, dqn), capsnet, distillation, ... 🧠项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations
如果你已经拿到annotated_deep_learning_paper_implementations仓库,想用它内置的 Stable Diffusion 实现把一段文本提示(prompt)渲染成图片,可以直接使用 text_to_image.py 这个脚本。该脚本基于仓库中的 Latent Diffusion 实现:加载预训练权重,用 CLIP 文本嵌入对扩散过程做条件引导,在潜空间采样后通过 autoencoder 解码出图像,并把结果保存为本地图片文件。整个实现不含训练代码,只跑推理流程;脚本会自动检测 CUDA,有 GPU 时用cuda:0,否则回退到cpu。
安装依赖并安装仓库代码
requirements.txt 声明了依赖版本下限,包括torch>=1.10、torchvision>=0.11、labml>=0.4.147、labml-helpers>=0.4.84、numpy>=1.19、Pillow>=6.2.1等。
在仓库根目录执行(对应 Makefile 中的install目标):
pip install -e .安装后labml_nn包可被导入,脚本中的from labml import lab、from labml_nn.diffusion.stable_diffusion...等导入才能解析。
准备模型权重
脚本把检查点路径写死在main()中:
txt2img = Txt2Img(checkpoint_path=lab.get_data_path() / 'stable-diffusion' / 'sd-v1-4.ckpt', sampler_name=opt.sampler_name, n_steps=opt.steps)也就是说,权重文件必须放在lab.get_data_path()返回的目录下、子目录stable-diffusion/中,且文件名为sd-v1-4.ckpt;命令行没有参数可以改这个路径。lab.get_data_path()来自labml包(requirements 中要求labml>=0.4.147),其具体落点由该包决定,运行前需要确认stable-diffusion/sd-v1-4.ckpt已经就位于此,否则脚本无法完成加载。
关于权重来源,Stable Diffusion 文档首页 说明:该实现基于官方 stable diffusion 仓库(CompVis/stable-diffusion),模型结构与官方保持一致,因此开源权重可以直接加载。项目文档没有给出权重文件的下载地址,需要读者自行获取与结构匹配的开源权重文件。
运行文本到图像脚本
从仓库根目录执行模块:
python -m labml_nn.diffusion.stable_diffusion.scripts.text_to_image --prompt "a painting of a virus monster playing guitar"--prompt后面的文本就是要渲染的提示词,上面的示例词正是脚本的默认值。脚本main()开头会调用set_seed(42)(见 util.py 中的set_seed,它同时固定random、numpy和torch的随机种子),所以同一台机器上重复运行同一命令,随机起点是确定的。
执行过程中,load_model 会按顺序打印这些阶段:Initialize autoencoder、Initialize CLIP Embedder、Initialize U-Net、Initialize Latent Diffusion model、Loading model from {path}、Load state,随后进入Generate阶段做采样与解码。看到这些阶段依次推进,说明模型已按 util.py 中定义的参数(如n_steps=1000、latent_scaling_factor=0.18215)构建并加载了状态字典。
命令行参数说明
脚本的 CLI 部分 定义了以下参数:
| 参数 | 默认值 | 用途(以文档说明为准) |
|---|---|---|
--prompt | a painting of a virus monster playing guitar | 要渲染的提示词(the prompt to render) |
--batch_size | 4 | 一批生成的图片数量 |
--sampler | ddim | 采样器,可选ddim/ddpm |
--steps | 50 | 采样步数(number of sampling steps) |
--scale | 7.5 | 无条件引导尺度,文档给出的公式为eps = eps(x, empty) + scale * (eps(x, cond) - eps(x, empty)) |
--flash | 未启用 | 是否使用 flash attention(store_true 开关) |
一个带完整参数的示例:
python -m labml_nn.diffusion.stable_diffusion.scripts.text_to_image \ --prompt "a painting of a virus monster playing guitar" \ --batch_size 4 \ --sampler ddim \ --steps 50 \ --scale 7.5参数取值只需替换成你自己的提示词与数量,命令其余部分可直接复制执行。
查看生成结果
生成结束后,脚本会把图片写入outputs目录(main()中dest_path='outputs'写死,不能通过命令行更改)。save_images 的文件命名规则是txt_前缀加 5 位编号、jpeg格式,因此默认--batch_size 4时,outputs/下会得到:
txt_00000.jpeg txt_00001.jpeg txt_00002.jpeg txt_00003.jpeg检查outputs/目录中是否出现与--batch_size数量一致、以txt_开头的.jpeg文件,是判断这次运行是否完成的直接依据。
可选调整与已记录的边界
- 切换采样器:
--sampler ddpm使用 DDPM 采样器,默认ddim使用 DDIM 采样器。从脚本代码看,--steps只通过n_steps传给了DDIMSampler;选ddpm时DDPMSampler(self.model)不接收该参数,步数由模型初始化时的n_steps=1000决定。 - 引导尺度:
--scale即文档中的无条件引导尺度 s。脚本逻辑显示,当uncond_scale == 1.0时不会计算空提示词的条件嵌入(un_cond = None),即关闭无条件引导;默认 7.5。 - flash attention:加
--flash会设置CrossAttention.use_flash_attention。Stable Diffusion 文档首页提到该实现可选集成 Flash Attention,在 RTX A6000 GPU 上可带来接近 50% 的性能提升——这是文档给出的唯一性能数据,且限定了 GPU 型号。 - 图像尺寸:
__call__中h、w默认 512,潜空间按 8 倍下采样(c=4、f=8),但 CLI 没有暴露尺寸参数,本脚本固定生成 512×512 图像。 - 不含训练:文档首页明确 “Our implementation does not contain training code”,本流程只做推理,不要期待在此实现里训练或微调模型。
如果outputs/目录没有出现预期的txt_*.jpeg文件,先回到“准备模型权重”一节确认stable-diffusion/sd-v1-4.ckpt是否位于lab.get_data_path()下的正确位置,这是脚本中唯一无法通过命令行绕过的前置条件。
【免费下载链接】annotated_deep_learning_paper_implementations🧑🏫 60+ Implementations/tutorials of deep learning papers with side-by-side notes 📝; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), 🎮 reinforcement learning (ppo, dqn), capsnet, distillation, ... 🧠项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考