Files
hair/hairline/comfyui.py
T
xsl ae5a3b8f1d refactor(comfyui): 模型切换时同步VAE + 完善编码器映射
- 新增 _vae_for_unet() 按模型自动切换 VAE(Z-Image用ae, Flux.2用flux2-vae)
- _clip_for_unet 增加 z-image 识别
- 为将来多模型切换做准备(当前 Q4 不受影响)
2026-07-26 16:12:43 +08:00

245 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""ComfyUI 客户端:用 add_hair.json / add_hair2.json 工作流跑生发图(Flux-2 inpaint)。
worker 不跑 Flux,只把「划线图 + 遮罩」的 RGBA 上传到远端 ComfyUI
(默认 http://10.60.74.221:8188,可用环境变量 COMFYUI_URL 覆盖),
替换工作流节点 26 的输入图、随机 seed,提交 /prompt,轮询 /history,取回 /view 输出。
ComfyUI 若开启了 HTTP Basic Authuser `admin` + 密码),所有请求都带凭据。
支持多工作流:run() 可通过 workflow_path 指定不同工作流 JSON,自动检测 SaveImage 输出节点。
"""
from __future__ import annotations
import copy
import json
import logging
import os
import random
import time
import uuid
import httpx
COMFYUI_URL = os.getenv("COMFYUI_URL", "http://127.0.0.1:8188").rstrip("/")
_WORKFLOW_DEFAULT = os.getenv(
"ADD_HAIR_WORKFLOW",
os.path.join(os.path.dirname(os.path.dirname(__file__)), "add_hair.json"),
)
COMFY_TIMEOUT = float(os.getenv("COMFYUI_TIMEOUT", "600")) # 单张出图最长等待(秒)
_REPO = os.path.dirname(os.path.dirname(__file__))
_INPUT_NODE = "26" # LoadImage:外部输入图(含 alpha 遮罩)
_SEED_NODE = "6" # RandomNoise
_PROMPT_NODE = "60" # JjkText:提示词
_UNET_NODE = "16" # UNETLoader / UnetLoaderGGUFFlux 模型加载
_CLIP_NODE = "61" # CLIPLoaderqwen 文本编码器
_VAE_NODE = "3" # VAELoader
# Flux 模型 → 配套文本编码器映射。切换 unet 时自动同步编码器,避免维度不匹配。
def _clip_for_unet(unet_name: str) -> str | None:
"""根据 unet 文件名推断配套的文本编码器文件名;无法推断返回 None。"""
low = unet_name.lower()
if "9b" in low: # Flux.2 9B 系列
return "qwen_3_8b_fp8mixed.safetensors"
if "z-image" in low: # Z-Image-Turbo 用 4B 编码器
return "qwen_3_4b.safetensors"
if "4b" in low: # Flux.2 4B
return "qwen_3_4b.safetensors"
return None
# Flux 模型 → 配套 VAE 映射。Z-Image 用 ae.safetensorsFlux.2 系列用 flux2-vae。
def _vae_for_unet(unet_name: str) -> str | None:
"""根据 unet 文件名推断配套 VAE 文件名;无法推断返回 None(保持工作流原值)。"""
low = unet_name.lower()
if "z-image" in low:
return "ae.safetensors"
return None # Flux.2 系列 vae 在工作流里已正确配置,不覆盖
_wf_cache: dict[str, dict] = {} # path → workflow JSON
_wf_output_node: dict[str, str] = {} # path → SaveImage 节点 ID
def _comfy_auth():
"""ComfyUI Basic Auth 凭据 (user, password)。
user:环境变量 COMFYUI_USER,默认 admin。
password:环境变量 COMFYUI_PASSWORD → worker_config.json.comfyui_password → password.txt。
无密码则返回 None(不带鉴权,兼容未开启 auth 的实例)。
"""
user = os.getenv("COMFYUI_USER", "admin")
pw = os.getenv("COMFYUI_PASSWORD")
if not pw:
cfg = os.path.join(_REPO, "worker_config.json")
if os.path.isfile(cfg):
try:
with open(cfg, encoding="utf-8") as f:
pw = json.load(f).get("comfyui_password")
except Exception: # noqa: BLE001
pw = None
if not pw:
pwfile = os.path.join(_REPO, "password.txt")
if os.path.isfile(pwfile):
with open(pwfile, encoding="utf-8") as f:
pw = f.read().strip()
return (user, pw) if pw else None
def _load_workflow(workflow_path: str | None = None) -> dict:
"""加载工作流 JSON(按路径缓存)。自动检测 SaveImage 节点 ID。"""
path = workflow_path or _WORKFLOW_DEFAULT
if path not in _wf_cache:
with open(path, encoding="utf-8") as f:
wf = json.load(f)
_wf_cache[path] = wf
# 自动检测 SaveImage 输出节点
for node_id, node in wf.items():
if node.get("class_type") == "SaveImage":
_wf_output_node[path] = node_id
break
else:
raise ValueError(f"工作流 {path} 中未找到 SaveImage 节点")
return _wf_cache[path]
def _get_output_node(workflow_path: str | None = None) -> str:
"""返回指定工作流的 SaveImage 节点 ID。"""
path = workflow_path or _WORKFLOW_DEFAULT
if path not in _wf_output_node:
_load_workflow(path) # 触发检测
return _wf_output_node[path]
def run(rgba_png_bytes: bytes, timeout: float = COMFY_TIMEOUT, prompt: str = None,
workflow_path: str | None = None, front: bool = False,
unet_name: str | None = None) -> bytes:
"""提交一次生发任务,返回输出 PNG 字节。失败抛异常。
prompt:非 None 时替换工作流节点60(JjkText)的文本;None 时用工作流内置默认提示词。
workflow_path:工作流 JSON 路径,None 则用默认 add_hair.json。
frontTrue 时任务插到 ComfyUI 队列最前(server 端 "front" 字段,队列号取负)。
接口2 对时延敏感用 True,避免排在接口3/5 的批量任务后面;其余接口保持 False。
unet_name:非 None 时改写工作流里的模型加载节点(节点16),动态切换 Flux 模型。
.safetensors → 保持 UNETLoader 节点类型不变,只替换 unet_name;
.gguf → 自动把节点类型改成 UnetLoaderGGUF(需装 ComfyUI-GGUF 插件)。
None 时用工作流内置默认模型。
"""
path = workflow_path or _WORKFLOW_DEFAULT
output_node = _get_output_node(path)
client_id = uuid.uuid4().hex
with httpx.Client(base_url=COMFYUI_URL, timeout=30.0, auth=_comfy_auth()) as cli:
# 1. 上传输入图(含 alpha 遮罩)到 ComfyUI input 目录
fname = f"hair_{client_id}.png"
r = cli.post("/upload/image", files={"image": (fname, rgba_png_bytes, "image/png")},
data={"overwrite": "true", "type": "input"})
r.raise_for_status()
up = r.json()
name = (up.get("subfolder") + "/" if up.get("subfolder") else "") + up["name"]
# 2. 改工作流:节点26 输入图 + 随机 seed
try:
import io as _io
from PIL import Image as _Img
_sz = _Img.open(_io.BytesIO(rgba_png_bytes)).size
logging.getLogger("hair.worker").info(
"ComfyUI 输入尺寸 %dx%d workflow=%s", _sz[0], _sz[1], os.path.basename(path))
except Exception: # noqa: BLE001
pass
wf = copy.deepcopy(_load_workflow(path))
wf[_INPUT_NODE]["inputs"]["image"] = name
wf[_SEED_NODE]["inputs"]["noise_seed"] = random.randint(0, 2**63 - 1)
if prompt is not None:
wf[_PROMPT_NODE]["inputs"]["text"] = prompt
if unet_name is not None:
node = wf.get(_UNET_NODE)
if node is not None:
# .gguf 需切换到 ComfyUI-GGUF 插件的 UnetLoaderGGUF 节点;
# .safetensors/.ckpt 保持原 UNETLoader 节点类型不变
if unet_name.lower().endswith(".gguf"):
node["class_type"] = "UnetLoaderGGUF"
else:
node["class_type"] = "UNETLoader"
node["inputs"]["unet_name"] = unet_name
# 同步切换配套文本编码器(4b→qwen_3_4b, 9b→qwen_3_8b),避免维度不匹配
clip_node = wf.get(_CLIP_NODE)
clip_name = _clip_for_unet(unet_name)
if clip_node is not None and clip_name is not None:
clip_node["inputs"]["clip_name"] = clip_name
# 同步切换 VAEZ-Image 用 ae.safetensorsFlux.2 保持 flux2-vae
vae_node = wf.get(_VAE_NODE)
vae_name = _vae_for_unet(unet_name)
if vae_node is not None and vae_name is not None:
vae_node["inputs"]["vae_name"] = vae_name
# 诊断:落盘实际提交的工作流 + 输入图,便于和手动 ComfyUI 跑的对比
try:
import os as _os
_diag = _os.path.join(_os.path.dirname(_os.path.dirname(_os.path.abspath(__file__))),
"log", "comfyui_last_submit")
_os.makedirs(_diag, exist_ok=True)
with open(_os.path.join(_diag, "workflow.json"), "w", encoding="utf-8") as _f:
json.dump(wf, _f, ensure_ascii=False, indent=2)
with open(_os.path.join(_diag, "input.png"), "wb") as _f:
_f.write(rgba_png_bytes)
with open(_os.path.join(_diag, "prompt.txt"), "w", encoding="utf-8") as _f:
_f.write(prompt if prompt is not None else "(None=用工作流内置默认)")
except Exception: # noqa: BLE001
pass
# 3. 提交(front=True 时插队到队列最前)
payload = {"prompt": wf, "client_id": client_id}
if front:
payload["front"] = True
r = cli.post("/prompt", json=payload)
r.raise_for_status()
prompt_id = r.json()["prompt_id"]
# 4. 轮询 /history
deadline = time.time() + timeout
outputs = None
while time.time() < deadline:
hr = cli.get(f"/history/{prompt_id}")
hr.raise_for_status()
hist = hr.json()
if prompt_id in hist:
entry = hist[prompt_id]
status = entry.get("status", {})
if status.get("status_str") == "error":
raise RuntimeError(f"ComfyUI 执行报错: {status}")
outputs = entry.get("outputs")
if outputs and output_node in outputs:
break
time.sleep(0.05)
if not outputs or output_node not in outputs:
raise TimeoutError(f"ComfyUI 出图超时({timeout}s) prompt_id={prompt_id}")
# 5. 取回输出图
imgs = outputs[output_node].get("images") or []
if not imgs:
raise RuntimeError("ComfyUI 输出无图像")
info = imgs[0]
vr = cli.get("/view", params={"filename": info["filename"],
"subfolder": info.get("subfolder", ""),
"type": info.get("type", "output")})
vr.raise_for_status()
return vr.content
def ping() -> bool:
"""探测 ComfyUI 是否在线(/system_stats)。"""
try:
with httpx.Client(base_url=COMFYUI_URL, timeout=3.0, auth=_comfy_auth()) as cli:
return cli.get("/system_stats").status_code == 200
except Exception: # noqa: BLE001
return False
if __name__ == "__main__":
import sys
inp = sys.argv[1] if len(sys.argv) > 1 else "tests/output/comfy_input.png"
print("ComfyUI:", COMFYUI_URL, "online:", ping())
with open(inp, "rb") as f:
png = run(f.read())
out = "tests/output/grown.png"
with open(out, "wb") as f:
f.write(png)
print(f"生发图已存 {out}{len(png)} bytes")