Spaces:
Running on Zero
Running on Zero
File size: 6,516 Bytes
ca9a89c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | """图片顺序、模型参考图能力和提示词编号的共享契约;不下载或改写图片。"""
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 []
|