Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
webui.py523 linesDownload Raw Back to root
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