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 []