Spaces:
Paused
Paused
Download caption.py from Quantumbraid/krea2train: direct link, hf CLI and curl.
- Browser
- Download file 8.52 kB
-
https://huggingface.co/spaces/Quantumbraid/krea2train/resolve/main/caption.py
- Command line
-
hf download hf://spaces/Quantumbraid/krea2train/caption.py
-
curl -L -o caption.py https://huggingface.co/spaces/Quantumbraid/krea2train/resolve/main/caption.py
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) | |