"""图片顺序、模型参考图能力和提示词编号的共享契约;不下载或改写图片。""" 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"(?:(? 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 []