File size: 8,515 Bytes
623fea3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
834a9ff
 
 
 
 
 
 
 
 
 
623fea3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
834a9ff
 
 
 
 
 
 
623fea3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
"""AI captioning + trigger suggestion for the Krea 2 LoRA trainer Space.

Runs on the Space itself (cpu-basic) by calling the HF Inference API for a multimodal LLM
(`google/gemma-4-31B-it`, served with vision by the **novita** provider — the default `auto`
route lands on an endpoint that returns empty text, so the provider is pinned).

The captioning token (`CAPTION_HF_TOKEN` secret) only ever calls the Inference API. It is
independent of the user's OAuth token (push/dataset) and the gated `KREA_TOKEN` (Krea weights).

Caption recipe follows the Krea 2 authors' training guidance:
  * STYLE LoRA  — describe only the *content* (subjects, poses, layout, setting), never the
    medium/technique/palette, then append the style trigger phrase (e.g. ", heavy impasto style").
  * OBJECT/CHARACTER LoRA — describe the scene with the subject referred to by its class noun,
    then append a unique trigger token (e.g. " b3@rcup").
"""

from __future__ import annotations

import base64
import io
import os

from huggingface_hub import InferenceClient
from PIL import Image

CAPTION_MODEL = "google/gemma-4-31B-it"
CAPTION_PROVIDER = "novita"
_MAX_SIDE = 768  # downscale before upload to keep the request small / fast

# Few-shot exemplars taken from the authors' reference captions (content only — the trigger is
# appended programmatically, so the examples here deliberately omit the trailing trigger).
_STYLE_EXAMPLES = [
    "A person is running forward in profile. The figure leans into the motion with their head "
    "tilted slightly down and long hair trailing horizontally behind. The arms are bent at the "
    "elbows, with one arm swung forward and the other pulled back toward the hip. One leg is "
    "extended backward, capturing a mid-stride movement. The figure is positioned centrally in a "
    "void of plain white.",
    "A fishing boat is stationed in a narrow canal between rows of multi-story buildings. The boat "
    "features a central cabin with windows and two vertical masts extending upwards. A red buoy "
    "hangs from the side of the hull. The water in the canal occupies the lower portion of the "
    "scene, while the sky is visible above the rooflines of the structures.",
    "A black sports car is positioned in the center of a wet road through a dense forest. The car "
    "faces forward, with its round headlights visible. The road surface is covered in puddles that "
    "reflect the front end. Tall coniferous trees line both sides of the road and a dense fog fills "
    "the space between the trees behind the vehicle.",
]
_OBJECT_EXAMPLES = [
    "A cup, sitting on a grainy wooden table with a grey door in the background. An iron stand has "
    "grey and black plastic containers in separate piles.",
    "A cup being held by a woman in her hand in the outdoors. The background is a textured patch of "
    "lawn grass.",
]


def _token() -> str:
    tok = os.environ.get("CAPTION_HF_TOKEN") or os.environ.get("HF_TOKEN") or ""
    if not tok:
        raise RuntimeError("AI captioning is unavailable: the CAPTION_HF_TOKEN secret is not set.")
    return tok


def _client() -> InferenceClient:
    return InferenceClient(model=CAPTION_MODEL, provider=CAPTION_PROVIDER, token=_token())


def _data_url(path: str) -> str:
    img = Image.open(path).convert("RGB")
    img.thumbnail((_MAX_SIDE, _MAX_SIDE))
    buf = io.BytesIO()
    img.save(buf, "JPEG", quality=90)
    return "data:image/jpeg;base64," + base64.b64encode(buf.getvalue()).decode()


def _ask(instruction: str, image_paths: list[str], max_tokens: int = 320,
         temperature: float = 0.4) -> str:
    content: list[dict] = [{"type": "text", "text": instruction}]
    for p in image_paths:
        content.append({"type": "image_url", "image_url": {"url": _data_url(p)}})
    r = _client().chat_completion(
        messages=[{"role": "user", "content": content}],
        max_tokens=max_tokens, temperature=temperature,
    )
    return (r.choices[0].message.content or "").strip()


def _clean(text: str) -> str:
    """Strip wrapping quotes / a leading 'Caption:' the model sometimes adds."""
    t = text.strip().strip('"').strip("'").strip()
    for prefix in ("Caption:", "caption:", "Trigger:", "trigger:"):
        if t.startswith(prefix):
            t = t[len(prefix):].strip()
    return t.rstrip()


def caption_one(image_path: str, concept_type: str, trigger: str) -> str:
    """Caption a single image for the given concept type, appending the trigger."""
    trigger = (trigger or "").strip()
    if concept_type == "custom":
        instruction = (
            "Write a concise, natural training caption that describes this image as you see it: "
            "the subjects, what they are doing, the setting, and the overall look. Write 1-3 plain "
            "declarative sentences. Return only the caption, with no preamble, labels or quotes."
        )
        cap = _clean(_ask(instruction, [image_path]))
        if trigger:
            cap = f"{trigger}, {cap}" if cap else trigger
        return cap
    if concept_type == "style":
        instruction = (
            "You are writing a training caption for a STYLE LoRA. Describe ONLY the literal "
            "content of the image: the subjects, their poses and actions, the key objects, their "
            "spatial arrangement, and the setting or background. Write 2-4 plain declarative "
            "sentences. Do NOT mention the artistic style, medium, technique, brushwork, lighting "
            "mood or palette, and do NOT use words like painting, illustration, drawing, render, "
            "sketch or photo. Match the tone of these examples:\n\n"
            + "\n\n".join(_STYLE_EXAMPLES)
            + "\n\nReturn only the caption sentence(s), with no preamble, labels or quotes."
        )
        cap = _clean(_ask(instruction, [image_path]))
        if trigger:
            cap = f"{cap.rstrip('.')}, {trigger}" if cap else trigger
        return cap
    # object / character
    instruction = (
        "You are writing a training caption for a LoRA of one specific subject. Describe the "
        "scene: where the subject is, what it is doing or how it is positioned, and the background "
        "or setting. Refer to the subject by its generic class noun (e.g. 'a cup', 'a dog'), never "
        "by a name. Write 1-3 plain declarative sentences. Match the tone of these examples:\n\n"
        + "\n\n".join(_OBJECT_EXAMPLES)
        + "\n\nReturn only the caption sentence(s), with no preamble, labels or quotes."
    )
    cap = _clean(_ask(instruction, [image_path]))
    if trigger:
        cap = f"{cap} {trigger}" if cap else trigger
    return cap


def suggest_trigger(image_paths: list[str], concept_type: str) -> str:
    """Suggest a trigger from 2-3 sample images: a style phrase, or a unique object token."""
    sample = list(image_paths)[:3]
    if not sample:
        raise gr_error("Upload images first.")
    if concept_type == "custom":
        instruction = (
            "Propose a SHORT unique trigger token for the concept shown in these images: a rare "
            "made-up token, optionally followed by a class noun. Examples: 'TOK', 'b3@rcup', "
            "'zxy style'. Return only the trigger, with no quotes or explanation."
        )
        return _clean(_ask(instruction, sample, max_tokens=16, temperature=0.7))
    if concept_type == "style":
        instruction = (
            "These images share one artistic style. Propose a SHORT distinctive trigger phrase "
            "naming that style: 2 to 5 words, ending with the word 'style'. Examples: 'heavy "
            "impasto style', 'monochrome ink wash style', 'flat pastel vector style'. Return only "
            "the phrase in lowercase, with no quotes or explanation."
        )
        return _clean(_ask(instruction, sample, max_tokens=24, temperature=0.6)).lower()
    instruction = (
        "These images show one specific subject. Propose a SHORT unique trigger for it: a rare "
        "made-up token, optionally followed by its class noun. Examples: 'b3@rcup', 'sks dog', "
        "'zxy sneaker'. Return only the trigger, with no quotes or explanation."
    )
    return _clean(_ask(instruction, sample, max_tokens=16, temperature=0.7))


def gr_error(msg: str):  # tiny indirection so this module stays importable without gradio
    try:
        import gradio as gr  # noqa: PLC0415
        return gr.Error(msg)
    except Exception:  # noqa: BLE001
        return ValueError(msg)