You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在StableDiffusionPipeline生成过程中获取迭代预览图像?

获取Stable Diffusion生成过程的中间预览(Discord机器人场景)

需求描述

用Hugging Face diffusers库的StableDiffusionPipeline开发Discord机器人,和朋友一起生成AI图像。希望在图像生成完成前获取预览,比如在20秒的生成过程中,保存每一次迭代(或每隔几秒)的图像,查看从模糊到清晰的演进过程。现有代码仅返回最终生成的图像,需要实现思路和技术提示。

现有基础代码:

class ImageGenerator:
    def __init__(self, socket_listener, pretty_logger, prisma):
        self.model = StableDiffusionPipeline.from_pretrained("runwayml/stable-diffusion-v1-5", revision="fp16", torch_dtype=torch.float16, use_auth_token=os.environ.get("HF_AUTH_TOKEN"))
        self.model = self.model.to("cuda")

    async def generate_image(self, data):
        start_time = time.time()
        with autocast("cuda"):
            image = self.model(data.description, height=self.default_height, width=self.default_width,
                               num_inference_steps=self.default_inference_steps, guidance_scale=self.default_guidance_scale)
            image.save(...)

实现思路与代码修改

核心方法:利用Pipeline的回调机制

StableDiffusionPipeline的__call__方法支持callback和callback_steps参数,通过自定义回调函数可以在指定步骤捕获中间生成的latent向量,再解码为图像保存。

具体步骤

  1. 定义回调函数
    回调函数会在每callback_steps步被调用,接收当前的latent、步骤数等参数,用pipeline的VAE组件解码latent为图像,然后保存或发送到Discord。

  2. 修改generate_image方法
    在调用model时传入回调函数和回调步长,比如每1步或每5步触发一次预览。

  3. 处理异步逻辑
    因为是Discord机器人的异步场景,回调里的保存/发送操作要避免阻塞,可结合asyncio处理。

修改后的完整代码示例

import torch
import time
from diffusers import StableDiffusionPipeline
from diffusers.pipelines.stable_diffusion import StableDiffusionPipelineOutput
from torch.cuda.amp import autocast
import os
from PIL import Image

class ImageGenerator:
    def __init__(self, socket_listener, pretty_logger, prisma):
        self.default_height = 512
        self.default_width = 512
        self.default_inference_steps = 50
        self.default_guidance_scale = 7.5
        self.model = StableDiffusionPipeline.from_pretrained(
            "runwayml/stable-diffusion-v1-5", 
            revision="fp16", 
            torch_dtype=torch.float16, 
            use_auth_token=os.environ.get("HF_AUTH_TOKEN")
        )
        self.model = self.model.to("cuda")
        self.logger = pretty_logger

    def save_intermediate_image(self, step, latents):
        # 将latent解码为图像
        with autocast("cuda"):
            # latent需要缩放回vae的输入范围(diffusers中latent默认缩放了0.18215倍)
            latents = 1 / 0.18215 * latents
            image = self.model.vae.decode(latents).sample
            # 归一化到0-1区间,再转成0-255的PIL图像
            image = (image / 2 + 0.5).clamp(0, 1)
            image = image.cpu().permute(0, 2, 3, 1).numpy()
            image = (image * 255).round().astype("uint8")
            pil_image = Image.fromarray(image[0])
            # 保存中间图,文件名包含步骤数
            pil_image.save(f"intermediate_step_{step}.png")
            self.logger.info(f"已保存中间步骤 {step} 的图像")

    def preview_callback(self, step: int, timestep: int, latents: torch.FloatTensor):
        # 每1步保存一次,也可以改成每N步触发,比如 step % 5 == 0
        self.save_intermediate_image(step, latents)

    async def generate_image(self, data):
        start_time = time.time()
        with autocast("cuda"):
            # 调用model时传入回调函数和回调步长
            output: StableDiffusionPipelineOutput = self.model(
                prompt=data.description,
                height=self.default_height,
                width=self.default_width,
                num_inference_steps=self.default_inference_steps,
                guidance_scale=self.default_guidance_scale,
                callback=self.preview_callback,
                callback_steps=1  # 可根据需求调整步长,比如设为5则每5步生成一次预览
            )
            # 保存最终图像
            output.images[0].save(f"final_image_{int(start_time)}.png")
            self.logger.info(f"生成完成,耗时 {time.time() - start_time:.2f} 秒")

关键细节说明

  • Latent缩放:diffusers中的latent是经过缩放的(乘以0.18215),解码前需要先缩放回去,否则图像会出现色彩/亮度异常。
  • 回调步长:callback_steps设为1会每步都生成预览,设为5则每5步生成一次,可根据生成速度和存储需求调整。
  • 异步适配:如果需要在回调里发送Discord消息,不能直接在同步回调中调用异步方法,可通过asyncio.run_coroutine_threadsafe将异步操作提交到事件循环。

内容的提问来源于stack exchange,提问作者jaal kamza

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.12 23:10:28