如何在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向量,再解码为图像保存。
具体步骤
定义回调函数
回调函数会在每callback_steps步被调用,接收当前的latent、步骤数等参数,用pipeline的VAE组件解码latent为图像,然后保存或发送到Discord。修改generate_image方法
在调用model时传入回调函数和回调步长,比如每1步或每5步触发一次预览。处理异步逻辑
因为是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
相关产品推荐
相关产品推荐

