Spaces:
Runtime error
Runtime error
File size: 3,936 Bytes
e4153e0 508c0b2 e4153e0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 | """
Couche d'orchestration partagee entre le endpoint WebSocket (FastAPI) et
l'interface Gradio, montes dans le MEME process (cf. decision d'architecture :
mount_gradio_app sur l'app FastAPI, pour que @spaces.GPU reste valide).
Contient :
- GPUQueue : serialise les appels GPU concurrents. Un seul GPU physique est
partage entre toutes les connexions - pas de vrai batching dynamique
possible sous ZeroGPU (decision documentee dans projet_2_avancement.md),
donc on implemente une file FIFO simple plutot qu'un faux "batching".
Un ThreadPoolExecutor a 1 worker traite les jobs strictement dans l'ordre
de soumission (FIFO garanti), en dehors de la boucle asyncio (les appels
modele sont bloquants/synchrones).
- run_inpaint_stream : pont entre le callback synchrone de diffusion (appele
depuis le thread GPU) et la boucle asyncio, sous forme de generateur async
consommable aussi bien par le handler WebSocket que par un callback Gradio
(Gradio accepte nativement les fonctions generatrices comme gestionnaires
d'evenements - meme mecanisme de streaming des deux cotes).
"""
import asyncio
from concurrent.futures import ThreadPoolExecutor
from typing import Any, AsyncGenerator, Callable
from PIL import Image
import pipeline
class GPUQueue:
def __init__(self):
self._executor = ThreadPoolExecutor(max_workers=1)
async def run(self, func: Callable, *args, **kwargs) -> Any:
loop = asyncio.get_running_loop()
return await loop.run_in_executor(self._executor, lambda: func(*args, **kwargs))
gpu_queue = GPUQueue()
async def run_locate(image: Image.Image, text_query: str) -> pipeline.LocateResult:
return await gpu_queue.run(pipeline.locate_and_segment, image, text_query)
async def run_inpaint_stream(
image: Image.Image,
dilated_mask: Image.Image,
prompt: str,
negative_prompt: str,
steps: int = 30,
guidance_scale: float = 7.5,
seed: int = 42,
) -> AsyncGenerator[tuple, None]:
"""Yield ("progress", step, total) pendant le debruitage, puis
("result", PIL.Image) une fois la generation terminee.
Sous ZeroGPU, PAS de step_callback : @spaces.GPU execute la fonction
decoree dans un process worker separe, et TOUS ses arguments doivent
etre picklables pour traverser cette frontiere IPC - une closure locale
(step_callback) ne l'est pas (`PicklingError: Can't pickle local object`,
rencontre au premier test reel sur le Space). Repli : pas de progression
incrementale sous ZeroGPU, seulement le resultat final."""
if pipeline.IS_ZERO_GPU:
result_image = await gpu_queue.run(
pipeline.generate_inpaint,
image, dilated_mask, prompt, negative_prompt,
steps, guidance_scale, seed, None,
)
yield ("result", result_image)
return
loop = asyncio.get_running_loop()
progress_queue: asyncio.Queue = asyncio.Queue()
def step_callback(step_index: int, total: int):
# Appele depuis le thread GPU (executor) : on ne peut pas toucher
# directement a une asyncio.Queue depuis un autre thread, on planifie
# donc le put() sur la boucle via call_soon_threadsafe.
loop.call_soon_threadsafe(progress_queue.put_nowait, (step_index, total))
task = asyncio.create_task(
gpu_queue.run(
pipeline.generate_inpaint,
image, dilated_mask, prompt, negative_prompt,
steps, guidance_scale, seed, step_callback,
)
)
while not task.done():
try:
step_index, total = await asyncio.wait_for(progress_queue.get(), timeout=0.2)
yield ("progress", step_index, total)
except asyncio.TimeoutError:
continue
result_image = await task
while not progress_queue.empty():
step_index, total = progress_queue.get_nowait()
yield ("progress", step_index, total)
yield ("result", result_image)
|