ImageGen-Studio / core /reference_inputs.py
BlueSkyXN's picture
Unify image references and guided UI/API workflows
ca9a89c verified
Raw History Blame Contribute Delete
6.52 kB
"""图片顺序、模型参考图能力和提示词编号的共享契约;不下载或改写图片。"""
from __future__ import annotations
import re
from PIL import Image
from core.model_capabilities import enabled_chains_for_model
from core.runtime_config import CONFIG
class ImageInputError(ValueError):
pass
# (pipeline input key, injector capacity)。应用上限仍受 RuntimeConfig 约束。
REFERENCE_CHAIN_SPECS = {
"qwen_image_edit": ("qwen_image_edit_data", 3),
"joyai_image": ("joyai_reference_data", 2),
"boogu_image_edit": ("boogu_edit_data", 10),
"reference_image": ("reference_image_data", 10),
"reference_latent": ("reference_latent_data", 10),
"hidream_o1_reference": ("hidream_o1_reference_data", 10),
"krea2_identity_edit": ("krea2_identity_edit_data", 2),
"krea2_style_reference": ("krea2_reference_edit_data", 3),
}
SOURCE_IMAGE_KEYS = {
"img2img": "img2img_image",
"inpaint": "inpaint_image",
"outpaint": "outpaint_image",
"hires_fix": "hires_image",
}
# 不匹配文件名/标识符;引号和反引号内的文字作为画面文字保留。
# 裸“图”不从构图、地图、草图等常用复合词中截取;保留完整的“参考图 / 图片”别名。
_REFERENCE = re.compile(
r"(?:(?<![A-Za-z_./\\-])(?:参考图|图片|(?<![构地蓝位视插草绘截贴])图)|(?<![A-Za-z0-9_./\\-])(?:img|image))[ \t]*([0-90-9]+)"
r"(?![0-90-9A-Za-z_./\\-])", re.IGNORECASE,
)
_LITERAL = re.compile(r'("[^"\n]*"|“[^”\n]*”|「[^」\n]*」|`[^`\n]*`)')
def reference_choices(model: str) -> dict[str, dict]:
enabled = enabled_chains_for_model(model)
available = [name for name in REFERENCE_CHAIN_SPECS if name in enabled]
roles = {}
for role, order in (
("auto", available),
("identity", ["krea2_identity_edit"]),
("style", ["krea2_style_reference"]),
):
chain = next((name for name in order if name in available), None)
if chain:
key, limit = REFERENCE_CHAIN_SPECS[chain]
roles[role] = {
"chain": chain, "input_key": key,
"max_images": min(limit, CONFIG.max_multi_images, CONFIG.max_reference_images),
}
return roles
def clear_native_references(values: dict) -> None:
for key, _ in REFERENCE_CHAIN_SPECS.values():
values[key] = []
def active_reference_groups(values: dict) -> dict[str, list]:
enabled = enabled_chains_for_model(values.get("model_display_name", ""))
return {
chain: [image for image in values.get(key, []) if image is not None]
for chain, (key, _) in REFERENCE_CHAIN_SPECS.items()
if chain in enabled and any(image is not None for image in values.get(key, []))
}
def image_bindings(count: int) -> list[dict]:
return [
{"id": f"img{i}", "index": i, "aliases": [f"图{i}", f"img{i}", f"image{i}"],
"model_reference": f"image {i}"}
for i in range(1, count + 1)
]
def validate_image_budget(images: list) -> None:
total = 0.0
for i, image in enumerate(images, 1):
if not isinstance(image, Image.Image):
continue # API 排队前可用未下载的 URL 校验结构。
megapixels = image.width * image.height / 1_000_000
if megapixels > CONFIG.max_input_megapixels:
raise ImageInputError(f"图{i} 为 {megapixels:.1f} MP,超过单图 {CONFIG.max_input_megapixels:g} MP 上限。")
total += megapixels
if total > CONFIG.max_reference_megapixels:
raise ImageInputError(f"图片累计 {total:.1f} MP,超过 {CONFIG.max_reference_megapixels:g} MP 上限;请减少或缩小图片。")
def normalize_image_references(prompt: str, count: int, *, independent=False, ambiguous=False) -> str:
"""仅统一显式编号;不是翻译器,也不承诺所有模型都懂位置指令。"""
def replace(match):
if independent:
raise ImageInputError("多图独立/图片×模型共用提示词,不能跨图引用编号;请写“当前图片”,或改用参考图编辑 / 融合。")
if ambiguous:
raise ImageInputError("多个原生参考链同时启用时图片编号不明确;请只保留一个参考链,或使用参考图编辑 / 融合入口。")
index = int(match.group(1))
if not 1 <= index <= count:
raise ImageInputError(f"提示词引用了 {match.group(0)},但本次只有 {count} 张有效图片;编号从 1 开始。")
separator = " " if _REFERENCE.match(match.string, match.end()) else ""
return f"image {index}{separator}"
parts = _LITERAL.split(prompt or "")
return "".join(part if index % 2 else _REFERENCE.sub(replace, part) for index, part in enumerate(parts))
def prepare_image_bindings(values: dict, *, independent=False) -> None:
"""在下载模型前处理当前有效输入;调用方持有自己的输入字典。"""
enabled = enabled_chains_for_model(values.get("model_display_name", ""))
for chain, (key, _) in REFERENCE_CHAIN_SPECS.items():
if chain not in enabled:
values[key] = []
groups = active_reference_groups(values)
task = values.get("task_type")
if task != "txt2img" and groups:
raise ImageInputError("原生参考图只能用于参考图编辑 / 融合(API 也兼容 txt2img + chain);普通重绘不要叠加参考链。")
for chain, images in groups.items():
limit = REFERENCE_CHAIN_SPECS[chain][1]
if len(images) > limit:
raise ImageInputError(f"{chain} 最多支持 {limit} 张参考图,当前为 {len(images)} 张;不会截断图片。")
count = sum(map(len, groups.values()))
if count > CONFIG.max_reference_images:
raise ImageInputError(f"本次参考图超过 {CONFIG.max_reference_images} 张上限。")
validate_image_budget([image for images in groups.values() for image in images])
if task in SOURCE_IMAGE_KEYS:
source = values.get(SOURCE_IMAGE_KEYS[task])
if task == "inpaint" and source is None:
source = (values.get("inpaint_image_dict") or {}).get("background")
count = int(source is not None)
for key in ("positive_prompt", "negative_prompt"):
values[key] = normalize_image_references(
values.get(key, ""), count, independent=independent, ambiguous=len(groups) > 1,
)
values["_image_references"] = image_bindings(count) if len(groups) <= 1 else []