krea2-lora-trainer / caption.py
multimodalart's picture
multimodalart HF Staff
Upload caption.py with huggingface_hub
834a9ff verified
Raw
History Blame Contribute Delete
8.52 kB
"""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)