wkea/blockdiffusion-api
0
1from __future__ import annotations
2
3"""
4BlockDiffusion WebUI(用于快速目测 checkpoint 的采样效果)
5
6说明:
7- 本文件只做“加载 ckpt -> 采样 -> 展示/保存”,复用项目现有扩散/CFG/EMA 逻辑。
8- 依赖:gradio(见 requirements.txt)。
9- 像素纹理预览:默认用 NEAREST 放大,方便看 16×16 的细节,不改变真实输出内容。
10"""
11
12import glob
13import math
14import os
15import re
16import time
17from dataclasses import dataclass
18from typing import Any, Dict, List, Optional, Tuple
19
20import torch
21from PIL import Image
22from torchvision.utils import make_grid
23
24from blockdiffusion.diffusion.gaussian_diffusion import GaussianDiffusion
25from blockdiffusion.diffusion.schedules import get_beta_schedule
26from blockdiffusion.models.text_unet16 import TextCondConfig, TextCondUNet16
27from blockdiffusion.models.unet16 import UNet16, UNet16Config
28from blockdiffusion.text.simple_tokenizer import SimpleVocab, tokenize
29from blockdiffusion.train.guidance import CFGuidanceModel
30from blockdiffusion.utils.checkpoint import load_checkpoint
31from blockdiffusion.utils.ema import EMA
32from blockdiffusion.utils.seed import get_device, seed_everything
33
34
35# ==============
36# Gradio import
37# ==============
38try:
39 import gradio as gr
40except Exception as e: # pragma: no cover
41 # 这里不要给“命令行安装指令”,只给明确缺依赖提示。
42 raise RuntimeError("缺少依赖:gradio。请先安装后再运行 WebUI。") from e
43
44
45@dataclass(frozen=True)
46class _Loaded:
47 """
48 缓存后的模型与配置(避免每次点击都重新加载 ckpt)。
49 注意:
50 - device 变化时会重新加载
51 - use_ema 变化时也会重新加载(因为 EMA 会 copy_to 覆盖权重)
52 """
53
54 key: Tuple[str, str, bool] # (ckpt_path, device_str, use_ema)
55 model_type: str
56 image_size: int
57 channels: int
58 timesteps: int
59 beta_schedule: str
60 unet_cfg: UNet16Config
61 text_cfg: Optional[TextCondConfig]
62 vocab: Optional[SimpleVocab]
63 model: torch.nn.Module
64 diffusion: GaussianDiffusion
65
66
67_CACHE: Optional[_Loaded] = None
68
69
70def _ensure_localhost_no_proxy() -> None:
71 """
72 规避 Windows/公司网络环境下的代理设置导致 localhost 访问被转发到代理,从而触发 Gradio 启动自检 502。
73 - gradio.launch() 内部会请求 `http://127.0.0.1:<port>/gradio_api/startup-events` 做自检
74 - 若 HTTP(S)_PROXY 存在且 NO_PROXY 未包含 localhost/127.0.0.1,可能出现 502
75 """
76
77 # 兼容大小写环境变量
78 existing = os.environ.get("NO_PROXY") or os.environ.get("no_proxy") or ""
79 parts = [p.strip() for p in existing.split(",") if p.strip()]
80 want = ["127.0.0.1", "localhost"]
81 for w in want:
82 if w not in parts:
83 parts.append(w)
84 val = ",".join(parts)
85 os.environ["NO_PROXY"] = val
86 os.environ["no_proxy"] = val
87
88
89def _parse_step_from_filename(path: str) -> int:
90 """
91 从 step_XXXX.pt 提取步数;提取失败则返回 -1,确保排序时落后。
92 """
93
94 base = os.path.basename(path)
95 m = re.search(r"step_(\d+)\.pt$", base)
96 if not m:
97 return -1
98 try:
99 return int(m.group(1))
100 except Exception:
101 return -1
102
103
104def list_checkpoints(outputs_dir: str = "outputs") -> List[str]:
105 """
106 扫描 outputs/**/checkpoints/step_*.pt,按 step 倒序排列(新的在前)。
107 """
108
109 pattern = os.path.join(outputs_dir, "**", "checkpoints", "step_*.pt")
110 paths = glob.glob(pattern, recursive=True)
111 paths.sort(key=lambda p: (_parse_step_from_filename(p), p), reverse=True)
112 return paths
113
114
115def _device_from_ui(device_text: str) -> torch.device:
116 """
117 UI 设备选择:
118 - "自动" => get_device(None)
119 - 其他 => get_device(str)
120 """
121
122 device_text = (device_text or "").strip()
123 if device_text in ("", "自动"):
124 return get_device(None)
125 return get_device(device_text)
126
127
128def _tensor_to_pil(x_chw_01: torch.Tensor) -> Image.Image:
129 """
130 (C,H,W) 且值域 [0,1] -> PIL(RGB/RGBA)
131 """
132
133 if x_chw_01.ndim != 3:
134 raise ValueError("x 必须是 (C,H,W)")
135 c = int(x_chw_01.shape[0])
136 if c not in (3, 4):
137 raise ValueError("仅支持 3(RGB) / 4(RGBA) 通道")
138
139 x = x_chw_01.clamp(0.0, 1.0)
140 arr = (x.permute(1, 2, 0) * 255.0).to(torch.uint8).cpu().numpy()
141 mode = "RGBA" if c == 4 else "RGB"
142 return Image.fromarray(arr, mode=mode)
143
144
145def _upscale_nearest(im: Image.Image, scale: int) -> Image.Image:
146 """
147 像素纹理预览:NEAREST 放大,不改像素风。
148 """
149
150 scale = int(scale)
151 if scale <= 1:
152 return im
153 w, h = im.size
154 return im.resize((w * scale, h * scale), resample=Image.NEAREST)
155
156
157@torch.no_grad()
158def _load_cached(ckpt_path: str, device_text: str, use_ema: bool) -> _Loaded:
159 """
160 加载/复用模型缓存。
161 """
162
163 global _CACHE
164 ckpt_path = os.path.abspath(ckpt_path)
165 key = (ckpt_path, (device_text or "").strip() or "自动", bool(use_ema))
166 if _CACHE is not None and _CACHE.key == key:
167 return _CACHE
168
169 device = _device_from_ui(device_text)
170 ckpt = load_checkpoint(ckpt_path, map_location="cpu")
171 meta = ckpt.get("meta", {}) or {}
172
173 model_type = str(meta.get("model_type", "uncond"))
174 unet_cfg_dict = meta.get("unet_cfg", None)
175 if unet_cfg_dict is None:
176 # 兼容:旧 ckpt 未保存配置
177 unet_cfg = UNet16Config()
178 timesteps = 1000
179 beta_schedule = "cosine"
180 image_size = 16
181 channels = 3
182 model_type = "uncond"
183 else:
184 unet_cfg = UNet16Config(**unet_cfg_dict)
185 diff_meta = meta.get("diffusion", {}) or {}
186 timesteps = int(diff_meta.get("timesteps", 1000))
187 beta_schedule = str(diff_meta.get("beta_schedule", "cosine"))
188 data_meta = meta.get("data", {}) or {}
189 image_size = int(data_meta.get("image_size", 16))
190 channels = int(data_meta.get("channels", unet_cfg.in_channels))
191
192 vocab = None
193 text_cfg = None
194 if model_type == "text_cond":
195 tok_dict = meta.get("tokenizer", None)
196 text_cfg_dict = meta.get("text_cfg", None)
197 if tok_dict is None or text_cfg_dict is None:
198 raise ValueError("该 checkpoint 标记为 text_cond,但缺少 tokenizer/text_cfg")
199 vocab = SimpleVocab.from_dict(tok_dict)
200 text_cfg = TextCondConfig(**text_cfg_dict)
201 model = TextCondUNet16(unet_cfg=unet_cfg, vocab=vocab, text_cfg=text_cfg).to(device)
202 else:
203 model = UNet16(unet_cfg).to(device)
204
205 model.load_state_dict(ckpt["model"], strict=True)
206 model.eval()
207
208 if use_ema and ("ema" in ckpt):
209 # 使用 EMA 权重采样(推荐),会覆盖 model 当前参数
210 ema = EMA(model, decay=0.9999)
211 ema.load_state_dict(ckpt["ema"])
212 ema.copy_to(model)
213 model.eval()
214
215 betas = get_beta_schedule(beta_schedule, timesteps).to(device)
216 diffusion = GaussianDiffusion(betas)
217
218 _CACHE = _Loaded(
219 key=key,
220 model_type=model_type,
221 image_size=image_size,
222 channels=channels,
223 timesteps=timesteps,
224 beta_schedule=beta_schedule,
225 unet_cfg=unet_cfg,
226 text_cfg=text_cfg,
227 vocab=vocab,
228 model=model,
229 diffusion=diffusion,
230 )
231 return _CACHE
232
233
234@torch.no_grad()
235def _sample(
236 ckpt_path: str,
237 prompt: str,
238 cfg_scale: float,
239 use_ema: bool,
240 use_ddim: bool,
241 ddim_steps: int,
242 ddim_eta: float,
243 n_images: int,
244 batch: int,
245 seed: int,
246 device_text: str,
247 save_dir: str,
248 preview_upscale: int,
249 gallery_n: int,
250) -> Tuple[str, Optional[str], Optional[List[Image.Image]], Optional[Image.Image]]:
251 """
252 返回:
253 - info 文本(用于 UI 展示)
254 - 保存路径(或 None)
255 - gallery 图片列表(或 None)
256 - 网格图 PIL(或 None)
257 """
258
259 ckpt_path = (ckpt_path or "").strip()
260 if not ckpt_path:
261 return "错误:请先选择/填写 checkpoint 路径。", None, None, None
262 if not os.path.isfile(ckpt_path):
263 return f"错误:checkpoint 不存在:{ckpt_path}", None, None, None
264
265 save_dir = (save_dir or "").strip() or "webui_samples"
266 os.makedirs(save_dir, exist_ok=True)
267
268 # 采样随机性(只影响本次点击)
269 seed_everything(int(seed), deterministic=False)
270
271 loaded = _load_cached(ckpt_path=ckpt_path, device_text=device_text, use_ema=use_ema)
272 device = _device_from_ui(device_text)
273
274 # 条件采样:仅 text_cond 支持 prompt + CFG
275 sampler_model: torch.nn.Module = loaded.model
276 prompt_tokens: List[str] = []
277 prompt_unk: List[str] = []
278 if (prompt or "").strip():
279 if loaded.model_type != "text_cond":
280 return "错误:当前 checkpoint 是无条件模型,不支持 prompt/CFG。", None, None, None
281 if loaded.vocab is None or loaded.text_cfg is None:
282 return "错误:条件模型缺少 tokenizer/text_cfg。", None, None, None
283
284 # 用于诊断:提示词是否大量变成 <unk>(这会导致“看起来不遵从提示词”)
285 prompt_tokens = tokenize(prompt)
286 for t in prompt_tokens:
287 tid = loaded.vocab.token_to_id.get(t, loaded.vocab.unk_id)
288 if int(tid) == int(loaded.vocab.unk_id):
289 prompt_unk.append(t)
290
291 tok = torch.tensor(
292 [loaded.vocab.encode(prompt, max_len=loaded.text_cfg.max_text_len)],
293 dtype=torch.long,
294 device=device,
295 )
296 m = tok != loaded.vocab.pad_id
297 cond_kwargs = {"tokens": tok, "mask": m}
298 un_tok = torch.full_like(tok, fill_value=loaded.vocab.pad_id)
299 un_m = torch.zeros_like(m)
300 uncond_kwargs = {"tokens": un_tok, "mask": un_m}
301 sampler_model = CFGuidanceModel(
302 base_model=loaded.model,
303 cond_kwargs=cond_kwargs,
304 uncond_kwargs=uncond_kwargs,
305 scale=float(cfg_scale),
306 ).to(device)
307 sampler_model.eval()
308 else:
309 # 不提供 prompt 时:对条件模型等价无条件采样
310 pass
311
312 n_images = int(max(1, n_images))
313 batch = int(max(1, batch))
314 gallery_n = int(max(0, min(gallery_n, n_images)))
315
316 # 分批采样
317 all_imgs: List[torch.Tensor] = []
318 remaining = n_images
319 t0 = time.time()
320 while remaining > 0:
321 b = min(batch, remaining)
322 shape = (b, loaded.channels, loaded.image_size, loaded.image_size)
323 if use_ddim:
324 x = loaded.diffusion.ddim_sample_loop(
325 sampler_model,
326 shape=shape,
327 device=device,
328 steps=int(ddim_steps),
329 eta=float(ddim_eta),
330 model_kwargs=None,
331 progress=False,
332 )
333 else:
334 x = loaded.diffusion.p_sample_loop(
335 sampler_model,
336 shape=shape,
337 device=device,
338 model_kwargs=None,
339 progress=False,
340 )
341 x = (x.clamp(-1, 1) + 1) * 0.5 # [0,1]
342 all_imgs.append(x.detach().cpu())
343 remaining -= b
344 dt = time.time() - t0
345
346 imgs = torch.cat(all_imgs, dim=0) # (N,C,H,W) on CPU
347
348 # 网格图(更适合一次看很多张)
349 nrow = int(max(1, math.isqrt(n_images)))
350 grid = make_grid(imgs, nrow=nrow, padding=1).clamp(0.0, 1.0) # (C,H,W)
351 grid_im_raw = _tensor_to_pil(grid) # 原始分辨率,用于保存
352 grid_im = _upscale_nearest(grid_im_raw, int(preview_upscale)) # 仅预览放大
353
354 # 单张 gallery(可选)
355 gallery: Optional[List[Image.Image]] = None
356 if gallery_n > 0:
357 gallery = []
358 for i in range(gallery_n):
359 im = _tensor_to_pil(imgs[i])
360 im = _upscale_nearest(im, int(preview_upscale))
361 gallery.append(im)
362
363 # 保存(默认保存网格图,文件名包含 step 与时间戳,方便对比)
364 step = _parse_step_from_filename(ckpt_path)
365 ts = time.strftime("%Y%m%d-%H%M%S")
366 out_name = f"grid_step_{step if step >= 0 else 'unk'}_{ts}.png"
367 out_path = os.path.join(save_dir, out_name)
368 grid_im_raw.save(out_path)
369
370 info = (
371 f"ckpt: {os.path.abspath(ckpt_path)}\n"
372 f"model_type: {loaded.model_type}\n"
373 f"size: {loaded.image_size}x{loaded.image_size} channels: {loaded.channels}\n"
374 f"timesteps: {loaded.timesteps} beta_schedule: {loaded.beta_schedule}\n"
375 f"device: {device}\n"
376 f"use_ema: {bool(use_ema)} sampler: {'DDIM' if use_ddim else 'DDPM'}\n"
377 f"n_images: {n_images} batch: {batch} seed: {int(seed)}\n"
378 f"prompt: {(prompt or '').strip() or '(none)'} cfg_scale: {float(cfg_scale):.2f}\n"
379 + (
380 f"prompt_tokens: {' '.join(prompt_tokens[:64])}\n"
381 f"prompt_unk({len(prompt_unk)}): {' '.join(prompt_unk[:64])}\n"
382 if prompt_tokens
383 else ""
384 )
385 + f"preview_upscale(NEAREST): {int(preview_upscale)}x\n"
386 + f"save: {os.path.abspath(out_path)}\n"
387 + f"time: {dt:.2f}s\n"
388 )
389 return info, os.path.abspath(out_path), gallery, grid_im
390
391
392def _build_ui() -> gr.Blocks:
393 ckpts = list_checkpoints("outputs")
394 default_ckpt = ckpts[0] if ckpts else ""
395
396 with gr.Blocks(title="BlockDiffusion WebUI") as demo:
397 gr.Markdown("### BlockDiffusion WebUI(采样预览)")
398
399 with gr.Row():
400 ckpt_quick = gr.Dropdown(
401 choices=ckpts,
402 value=default_ckpt or None,
403 label="快速选择 checkpoint(从 outputs 扫描)",
404 )
405 refresh = gr.Button("刷新列表")
406
407 ckpt_path = gr.Textbox(value=default_ckpt, label="Checkpoint 路径(可手填)")
408
409 with gr.Row():
410 prompt = gr.Textbox(value="", label="Prompt(条件模型可用;留空=无条件)")
411 cfg_scale = gr.Slider(minimum=0.0, maximum=12.0, value=3.0, step=0.1, label="CFG scale")
412
413 with gr.Row():
414 use_ema = gr.Checkbox(value=True, label="使用 EMA 权重(推荐)")
415 use_ddim = gr.Checkbox(value=True, label="使用 DDIM(更快)")
416 ddim_steps = gr.Slider(minimum=10, maximum=200, value=50, step=1, label="DDIM steps")
417 ddim_eta = gr.Slider(minimum=0.0, maximum=1.0, value=0.0, step=0.01, label="DDIM eta")
418
419 with gr.Row():
420 n_images = gr.Slider(minimum=1, maximum=256, value=64, step=1, label="生成张数 N")
421 batch = gr.Slider(minimum=1, maximum=256, value=64, step=1, label="采样 batch(显存不够就调小)")
422 seed = gr.Number(value=123, precision=0, label="随机种子 seed")
423
424 with gr.Row():
425 device_text = gr.Dropdown(
426 choices=["自动", "cuda", "cuda:0", "cpu"],
427 value="自动",
428 label="设备",
429 )
430 save_dir = gr.Textbox(value="webui_samples", label="保存目录")
431 preview_upscale = gr.Slider(minimum=1, maximum=64, value=16, step=1, label="预览放大倍数(NEAREST)")
432 gallery_n = gr.Slider(minimum=0, maximum=64, value=16, step=1, label="单张预览数量(0=关闭)")
433
434 run = gr.Button("生成")
435
436 info = gr.Textbox(label="信息/日志", lines=10)
437 grid_out = gr.Image(label="网格图(已放大预览)", type="pil")
438 gallery_out = gr.Gallery(label="单张预览(已放大)", columns=8, rows=2)
439 saved_path = gr.Textbox(label="已保存路径", lines=1)
440
441 def _on_refresh():
442 new_ckpts = list_checkpoints("outputs")
443 return gr.update(choices=new_ckpts, value=(new_ckpts[0] if new_ckpts else None))
444
445 def _on_pick(p: Optional[str]) -> str:
446 return p or ""
447
448 def _on_run(
449 ckpt_path_v: str,
450 prompt_v: str,
451 cfg_scale_v: float,
452 use_ema_v: bool,
453 use_ddim_v: bool,
454 ddim_steps_v: int,
455 ddim_eta_v: float,
456 n_images_v: int,
457 batch_v: int,
458 seed_v: int,
459 device_text_v: str,
460 save_dir_v: str,
461 preview_upscale_v: int,
462 gallery_n_v: int,
463 ):
464 try:
465 txt, out_path, gal, grid = _sample(
466 ckpt_path=ckpt_path_v,
467 prompt=prompt_v,
468 cfg_scale=float(cfg_scale_v),
469 use_ema=bool(use_ema_v),
470 use_ddim=bool(use_ddim_v),
471 ddim_steps=int(ddim_steps_v),
472 ddim_eta=float(ddim_eta_v),
473 n_images=int(n_images_v),
474 batch=int(batch_v),
475 seed=int(seed_v),
476 device_text=device_text_v,
477 save_dir=save_dir_v,
478 preview_upscale=int(preview_upscale_v),
479 gallery_n=int(gallery_n_v),
480 )
481 return txt, grid, (gal or []), (out_path or "")
482 except Exception as e:
483 # UI 侧直接给出错误信息,避免整个服务崩溃
484 return f"错误:{e}", None, [], ""
485
486 refresh.click(_on_refresh, inputs=[], outputs=[ckpt_quick])
487 ckpt_quick.change(_on_pick, inputs=[ckpt_quick], outputs=[ckpt_path])
488 run.click(
489 _on_run,
490 inputs=[
491 ckpt_path,
492 prompt,
493 cfg_scale,
494 use_ema,
495 use_ddim,
496 ddim_steps,
497 ddim_eta,
498 n_images,
499 batch,
500 seed,
501 device_text,
502 save_dir,
503 preview_upscale,
504 gallery_n,
505 ],
506 outputs=[info, grid_out, gallery_out, saved_path],
507 )
508
509 return demo
510
511
512def main() -> None:
513 demo = _build_ui()
514 # 默认监听 127.0.0.1:仅本机访问更安全
515 _ensure_localhost_no_proxy()
516 demo.launch(server_name="127.0.0.1", share=False)
517
518
519if __name__ == "__main__":
520 main()
521
522
523 