Spaces:
Paused
Paused
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)
|