Spaces:
Running on Zero
Running on Zero
Download core/reference_inputs.py from BlueSkyXN/ImageGen-Studio: direct link, hf CLI and curl.
- Browser
- Download file 6.52 kB
-
https://huggingface.co/spaces/BlueSkyXN/ImageGen-Studio/resolve/main/core/reference_inputs.py
- Command line
-
hf download hf://spaces/BlueSkyXN/ImageGen-Studio/core/reference_inputs.py
-
curl -L -o reference_inputs.py https://huggingface.co/spaces/BlueSkyXN/ImageGen-Studio/resolve/main/core/reference_inputs.py
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 [] | |