Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
modelscope_ui_app.py236 linesDownload Raw Back to root
1from __future__ import annotations
2
3"""
4ModelScope SwingDeploy(FC)UI 函数 app.py(仅适配 text-to-image-synthesis)。
5
6用法:
7- 将本文件内容复制到 FC 的 model_ui_func 里作为 app.py
8- 环境变量由部署自动注入:
9  - MODEL_ID / MODEL_VERSION / TASK / API_URL / MODEL_BACKEND
10
11说明:
12- UI 会把请求转发到 API_URL + "/invoke"(pipeline 模式默认);
13- 兼容服务端返回 image 的两种形式:
14  - base64 字符串(可带 data:image/... 前缀)
15  - HWC 的嵌套数组(uint8)
16"""
17
18import base64
19import io
20import os
21import json
22from typing import Any, Dict, List, Optional, Tuple, Union
23
24import gradio as gr
25import numpy as np
26import requests
27from PIL import Image
28
29MODEL_ID = os.getenv("MODEL_ID", "") or ""
30MODEL_VERSION = os.getenv("MODEL_VERSION", "") or ""
31TASK = os.getenv("TASK", "") or ""
32API_URL = (os.getenv("API_URL", "") or "").rstrip("/")
33MODEL_BACKEND = os.getenv("MODEL_BACKEND", "pipeline") or "pipeline"
34
35# pipeline 默认推理路径为 /invoke;ollama 会是 /v1(这里只做兼容,不建议用 UI 跑 ollama)
36INFER_PATH = "v1" if MODEL_BACKEND == "ollama" else "invoke"
37INFER_URL = (API_URL + "/" + INFER_PATH) if API_URL else ""
38
39
40def _unwrap_response(resp_json: Any) -> Dict[str, Any]:
41    """
42    SwingDeploy 常见返回:
43    - {"Code":200,"Data":{...},...}
44    也兼容直接返回 dict 的情况。
45    """
46    if isinstance(resp_json, dict):
47        if "Code" in resp_json and resp_json.get("Code") != 200:
48            raise gr.Error(f"推理失败:Code={resp_json.get('Code')} Message={resp_json.get('Message')}")
49        data = resp_json.get("Data", None)
50        if isinstance(data, dict):
51            return data
52        # 有些情况下 Data 不是 dict,就当作整体返回
53        return resp_json
54    raise gr.Error("推理返回不是JSON对象")
55
56
57def _decode_base64_image(s: str) -> Optional[Image.Image]:
58    s = (s or "").strip()
59    if not s:
60        return None
61    if s.startswith("data:image"):
62        parts = s.split(",", 1)
63        if len(parts) == 2:
64            s = parts[1].strip()
65    try:
66        raw = base64.b64decode(s, validate=False)
67        return Image.open(io.BytesIO(raw)).convert("RGBA")
68    except Exception:
69        return None
70
71
72def _decode_array_image(x: Any) -> Optional[Image.Image]:
73    try:
74        a = np.asarray(x)
75    except Exception:
76        return None
77    if a.ndim != 3 or a.shape[-1] not in (3, 4):
78        return None
79    if a.dtype != np.uint8:
80        try:
81            a = a.astype(np.uint8)
82        except Exception:
83            return None
84    mode = "RGBA" if a.shape[-1] == 4 else "RGB"
85    try:
86        return Image.fromarray(a, mode=mode)
87    except Exception:
88        return None
89
90
91def _extract_images(data: Dict[str, Any]) -> List[Image.Image]:
92    """
93    兼容 keys:
94    - output_img / output_imgs
95    - 其他上游包装层(递归查找)
96    """
97
98    def pick(obj: Any) -> List[Any]:
99        found: List[Any] = []
100        if isinstance(obj, dict):
101            for k in ("output_imgs", "output_img", "images", "image", "imgs", "img"):
102                if k in obj and obj[k] is not None:
103                    found.append(obj[k])
104            for v in obj.values():
105                if isinstance(v, (dict, list)):
106                    found.extend(pick(v))
107        elif isinstance(obj, list):
108            for v in obj:
109                if isinstance(v, (dict, list)):
110                    found.extend(pick(v))
111        return found
112
113    candidates = pick(data)
114    images: List[Image.Image] = []
115
116    def push_one(v: Any) -> None:
117        if v is None:
118            return
119        if isinstance(v, Image.Image):
120            images.append(v)
121            return
122        if isinstance(v, str):
123            im = _decode_base64_image(v)
124            if im is not None:
125                images.append(im)
126            return
127        # 数组(HWC)
128        im = _decode_array_image(v)
129        if im is not None:
130            images.append(im)
131
132    for c in candidates:
133        # output_imgs 可能是 list[img]
134        if isinstance(c, list):
135            # 也可能是单张 HWC 嵌套数组
136            im_single = _decode_array_image(c)
137            if im_single is not None:
138                images.append(im_single)
139                continue
140            for one in c:
141                push_one(one)
142        else:
143            push_one(c)
144
145    if not images:
146        raise gr.Error("推理成功但未在返回中找到图片(output_img/output_imgs)")
147    return images
148
149
150def _post_infer(payload: Dict[str, Any]) -> Dict[str, Any]:
151    if not INFER_URL:
152        raise gr.Error("缺少 API_URL 环境变量,无法调用推理接口")
153    try:
154        r = requests.post(INFER_URL, json=payload, timeout=600)
155    except Exception as e:
156        raise gr.Error(f"请求失败:{e!r}")
157    try:
158        j = r.json()
159    except Exception:
160        raise gr.Error(f"返回非JSON,HTTP={r.status_code}")
161    return _unwrap_response(j)
162
163
164def build_demo() -> gr.Blocks:
165    title = "ModelScope × FC:text-to-image-synthesis"
166    article = (
167        f"- 模型ID:{MODEL_ID}\n"
168        f"- 模型版本:{MODEL_VERSION}\n"
169        f"- TASK:{TASK}\n"
170        f"- 推理URL:{INFER_URL}\n"
171    )
172
173    with gr.Blocks(title=title) as demo:
174        gr.Markdown(f"## {title}")
175        gr.Markdown(article)
176
177        if TASK != "text-to-image-synthesis":
178            gr.Markdown("### 当前 UI 仅支持 TASK=text-to-image-synthesis")
179            return demo
180
181        with gr.Row():
182            prompt = gr.Textbox(label="Prompt", value="stone block texture", lines=2)
183
184        def build_seed_component():
185            # 兼容旧版 gradio:可能没有 gr.Number
186            if hasattr(gr, "Number"):
187                return gr.Number(label="seed(-1表示不传)", value=-1, precision=0)
188            return gr.Textbox(label="seed(-1表示不传)", value="-1")
189
190        def build_json_component():
191            # 兼容旧版 gradio:可能没有 gr.JSON
192            if hasattr(gr, "JSON"):
193                return gr.JSON(label="原始输出(Data)")
194            return gr.Textbox(label="原始输出(Data)", lines=12)
195
196        with gr.Row():
197            n = gr.Slider(label="n(生成张数)", minimum=1, maximum=8, step=1, value=1)
198            ddim_steps = gr.Slider(label="ddim_steps", minimum=5, maximum=100, step=1, value=50)
199            cfg_scale = gr.Slider(label="cfg_scale(仅 text_cond 生效)", minimum=0.0, maximum=20.0, step=0.5, value=3.0)
200            seed = build_seed_component()
201
202        run_btn = gr.Button("生成")
203        # 兼容旧版 gradio:Gallery 可能没有 .style()
204        gallery = gr.Gallery(label="结果")
205        raw_json = build_json_component()
206
207        def handler(p: str, n_v: int, steps_v: int, cfg_v: float, seed_v: Any):
208            payload: Dict[str, Any] = {
209                "input": {"text": p},
210                "parameters": {"n": int(n_v), "ddim_steps": int(steps_v), "cfg_scale": float(cfg_v)},
211            }
212            # seed 兼容:Number 会传 float/int,Textbox 会传 str
213            try:
214                seed_int = int(float(seed_v))
215            except Exception:
216                seed_int = -1
217            if seed_int >= 0:
218                payload["parameters"]["seed"] = seed_int
219            data = _post_infer(payload)
220            imgs = _extract_images(data)
221            # Gallery 支持 list[PIL.Image]
222            if hasattr(gr, "JSON"):
223                return imgs, data
224            return imgs, json.dumps(data, ensure_ascii=False, indent=2)
225
226        run_btn.click(handler, inputs=[prompt, n, ddim_steps, cfg_scale, seed], outputs=[gallery, raw_json])
227
228    return demo
229
230
231if __name__ == "__main__":
232    demo = build_demo()
233    demo.launch(server_name="0.0.0.0")
234
235
236