wkea/blockdiffusion-api
0
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 