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)