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

如何在同步图像生成流程中使用FastAPI WebSocket异步回调?

解决同步Stable Diffusion Pipeline回调中调用异步WebSocket的问题

错误根源

  1. Stable Diffusion的callback是同步函数,无法直接调用异步的websocket.send_bytes,直接调用会导致协程未被调度执行,出现RuntimeWarning: coroutine was never awaited。
  2. 错误地在functools.partial中使用await,await仅能在异步函数内部使用,不能直接作用于函数对象,引发TypeError。
  3. 同步的图像生成过程会阻塞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.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 01:04:59