如何在同步图像生成流程中使用FastAPI WebSocket异步回调?
解决同步Stable Diffusion Pipeline回调中调用异步WebSocket的问题
错误根源
- Stable Diffusion的
callback是同步函数,无法直接调用异步的websocket.send_bytes,直接调用会导致协程未被调度执行,出现RuntimeWarning: coroutine was never awaited。 - 错误地在
functools.partial中使用await,await仅能在异步函数内部使用,不能直接作用于函数对象,引发TypeError。 - 同步的图像生成过程会阻塞FastAPI的事件循环,导致WebSocket消息无法及时发送,甚至影响整个服务的响应性。
解决方案
1. 封装同步回调,调度异步任务
用同步函数作为SD的回调,在内部通过主线程的事件循环调度异步的WebSocket发送逻辑,确保异步任务能被正确执行。
2. 将同步生成任务移至线程池
把SD的图像生成过程放到线程池执行,避免阻塞FastAPI的事件循环,保证WebSocket消息和其他请求能正常处理。
修正后的完整代码
import asyncio import io import gc import traceback import functools import torch from fastapi import WebSocket, APIRouter router = APIRouter() # 假设stableDiffusionPipeline和mergeImageIfNecessary已提前初始化 # stableDiffusionPipeline = ... # mergeImageIfNecessary = ... @router.websocket("/ws") async def websocket_endpoint(websocket: WebSocket): await websocket.accept() prompt = await websocket.receive_text() if not prompt: await websocket.send_text("Error: Empty prompt") return # 获取当前事件循环(主线程的事件循环) loop = asyncio.get_running_loop() async def pipelineCallback(iteration, t, latents): print(f"Processing step {iteration}...") with torch.no_grad(): # 解码latents为图像 latents = 1 / 0.18215 * latents image = stableDiffusionPipeline.vae.decode(latents).sample image = (image / 2 + 0.5).clamp(0, 1) image = image.cpu().permute(0, 2, 3, 1).float().numpy() # 转换为PIL图像并处理 images = stableDiffusionPipeline.numpy_to_pil(image) image = mergeImageIfNecessary(images) # 转换为字节并发送 buffer = io.BytesIO() image.save(buffer, format="PNG") buffer.seek(0) await websocket.send_bytes(buffer.read()) def callbackCaller(iteration, t, latents): # 跨线程安全地调度异步任务(因为SD生成在线程池线程执行) def schedule_task(): asyncio.create_task(pipelineCallback(iteration, t, latents)) loop.call_soon_threadsafe(schedule_task) try: torch.Generator("cuda").manual_seed(0) # 将同步的SD生成任务放到线程池执行,避免阻塞事件循环 images = await loop.run_in_executor( None, lambda: stableDiffusionPipeline( prompt=prompt, width=512, height=512, num_inference_steps=100, guidance_scale=8, callback=callbackCaller, # 传递同步回调函数 callback_steps=10, num_images_per_prompt=1 ).images ) # 处理最终生成的图像 image = mergeImageIfNecessary(images) gc.collect() torch.cuda.empty_cache() buffer = io.BytesIO() image.save(buffer, format="PNG") buffer.seek(0) final_image_bytes = buffer.read() if not final_image_bytes: await websocket.send_text("Error: Image generation failed.") return await websocket.send_bytes(final_image_bytes) print("Closing socket...") await websocket.close() except Exception as e: traceback.print_exc() await websocket.send_text(f"Error: {str(e)}")
关键修改说明
- 回调调度:用
loop.call_soon_threadsafe确保在主线程的事件循环中调度异步任务,避免线程池线程中无事件循环的问题。 - 线程池执行生成:通过
loop.run_in_executor把SD的同步生成逻辑移至线程池,释放事件循环,保证WebSocket消息能及时发送。 - 移除错误的await用法:直接传递同步的
callbackCaller作为SD的回调,不再在functools.partial中误用await。
内容的提问来源于stack exchange,提问作者Julian S.
相关产品推荐
相关产品推荐

