"""Pluggable local captioners for LoRA dataset building. All captioners expose the same interface: ``caption(images: list[PIL.Image]) -> str``. ``images`` is a list so video callers can pass multiple sampled frames; single-image captioners just use the first one. - JoyCaptionCaptioner: fancyfeast/llama-joycaption-beta-one-hf-llava (AutoProcessor + LlavaForConditionalGeneration, confirmed against the model card). The exact prompt wording for a dedicated "training caption" mode is not documented on the model card, so the instruction text here is our own default and is overridable via --caption-prompt. - Florence2Captioner: microsoft/Florence-2-large, task. - WD14Captioner: SmilingWolf/wd-vit-tagger-v3 onnx tagger (booru-style tags), preprocessing confirmed against the reference wd-tagger Space source (white-pad-to-square, RGB->BGR, no normalization, category 9/4/0 = rating/character/general). """ import csv import numpy as np import torch from PIL import Image class JoyCaptionCaptioner: MODEL_ID = "fancyfeast/llama-joycaption-beta-one-hf-llava" DEFAULT_PROMPT = ( "Write a detailed but concise caption for this image, suitable for training " "an image/video LoRA. Describe the subject's appearance, pose, clothing, " "setting, and lighting in plain descriptive sentences. Do not guess names." ) def __init__(self, device="cuda", load_in_4bit=False, prompt=None, system_prompt=None): from transformers import AutoProcessor, LlavaForConditionalGeneration self.processor = AutoProcessor.from_pretrained(self.MODEL_ID) load_kwargs = {} if load_in_4bit: from transformers import BitsAndBytesConfig load_kwargs["quantization_config"] = BitsAndBytesConfig(load_in_4bit=True) load_kwargs["device_map"] = "auto" else: load_kwargs["torch_dtype"] = torch.bfloat16 load_kwargs["device_map"] = device self.model = LlavaForConditionalGeneration.from_pretrained(self.MODEL_ID, **load_kwargs) self.model.eval() self.system_prompt = system_prompt or "You are a helpful image captioner." self.prompt = prompt or self.DEFAULT_PROMPT def caption(self, images): image = images[0].convert("RGB") convo = [ {"role": "system", "content": self.system_prompt}, {"role": "user", "content": self.prompt}, ] convo_string = self.processor.apply_chat_template( convo, tokenize=False, add_generation_prompt=True ) inputs = self.processor(text=[convo_string], images=[image], return_tensors="pt") inputs = {k: v.to(self.model.device) for k, v in inputs.items()} if "pixel_values" in inputs: inputs["pixel_values"] = inputs["pixel_values"].to(self.model.dtype) with torch.no_grad(): generate_ids = self.model.generate( **inputs, max_new_tokens=300, do_sample=True, suppress_tokens=None, use_cache=True, temperature=0.6, top_k=None, top_p=0.9, )[0] generate_ids = generate_ids[inputs["input_ids"].shape[1]:] text = self.processor.tokenizer.decode( generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False ) return text.strip() class Florence2Captioner: MODEL_ID = "microsoft/Florence-2-large" TASK = "" def __init__(self, device="cuda"): from transformers import AutoModelForCausalLM, AutoProcessor dtype = torch.float16 if device.startswith("cuda") else torch.float32 self.device = device self.model = ( AutoModelForCausalLM.from_pretrained(self.MODEL_ID, trust_remote_code=True, torch_dtype=dtype) .to(device) .eval() ) self.processor = AutoProcessor.from_pretrained(self.MODEL_ID, trust_remote_code=True) def caption(self, images): image = images[0].convert("RGB") inputs = self.processor(text=self.TASK, images=image, return_tensors="pt") inputs = {k: v.to(self.device, self.model.dtype) if v.is_floating_point() else v.to(self.device) for k, v in inputs.items()} with torch.no_grad(): generated_ids = self.model.generate( input_ids=inputs["input_ids"], pixel_values=inputs["pixel_values"], max_new_tokens=200, num_beams=3, do_sample=False, ) text = self.processor.batch_decode(generated_ids, skip_special_tokens=False)[0] parsed = self.processor.post_process_generation( text, task=self.TASK, image_size=(image.width, image.height) ) return parsed[self.TASK].strip() class WD14Captioner: REPO_ID = "SmilingWolf/wd-vit-tagger-v3" MODEL_FILE = "model.onnx" TAGS_FILE = "selected_tags.csv" def __init__(self, general_thresh=0.35, character_thresh=0.75, use_gpu=True): import onnxruntime as ort from huggingface_hub import hf_hub_download model_path = hf_hub_download(self.REPO_ID, self.MODEL_FILE) tags_path = hf_hub_download(self.REPO_ID, self.TAGS_FILE) providers = ["CPUExecutionProvider"] if use_gpu: providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] self.session = ort.InferenceSession(model_path, providers=providers) self.input_name = self.session.get_inputs()[0].name self.output_name = self.session.get_outputs()[0].name self.target_size = int(self.session.get_inputs()[0].shape[1]) self.general_thresh = general_thresh self.character_thresh = character_thresh self.tag_names, self.general_idx, self.character_idx = self._load_tags(tags_path) @staticmethod def _load_tags(path): names, general, character = [], [], [] with open(path, newline="", encoding="utf-8") as f: for i, row in enumerate(csv.DictReader(f)): names.append(row["name"]) cat = int(row["category"]) if cat == 4: character.append(i) elif cat != 9: # skip rating (9); keep general (0) and any others general.append(i) return names, general, character def _preprocess(self, image: Image.Image): image = image.convert("RGBA") canvas = Image.new("RGBA", image.size, (255, 255, 255, 255)) canvas.alpha_composite(image) image = canvas.convert("RGB") w, h = image.size size = max(w, h) padded = Image.new("RGB", (size, size), (255, 255, 255)) padded.paste(image, ((size - w) // 2, (size - h) // 2)) padded = padded.resize((self.target_size, self.target_size), Image.BICUBIC) arr = np.asarray(padded, dtype=np.float32) arr = arr[:, :, ::-1] # RGB -> BGR return np.expand_dims(arr, axis=0) def caption(self, images): image = images[0] arr = self._preprocess(image) preds = self.session.run([self.output_name], {self.input_name: arr})[0][0] general = [(self.tag_names[i], preds[i]) for i in self.general_idx if preds[i] >= self.general_thresh] character = [(self.tag_names[i], preds[i]) for i in self.character_idx if preds[i] >= self.character_thresh] general.sort(key=lambda t: -t[1]) character.sort(key=lambda t: -t[1]) tags = [t.replace("_", " ") for t, _ in character] + [t.replace("_", " ") for t, _ in general] return ", ".join(tags) def build_captioner(args): name = args.captioner if name == "joycaption": return JoyCaptionCaptioner( device=args.device, load_in_4bit=args.load_in_4bit, prompt=args.caption_prompt, ) if name == "florence2": return Florence2Captioner(device=args.device) if name == "wd14": return WD14Captioner( general_thresh=args.wd14_general_thresh, character_thresh=args.wd14_character_thresh, use_gpu=args.device.startswith("cuda"), ) raise ValueError(f"Unknown captioner: {name}") def with_trigger(caption, trigger): caption = (caption or "").strip() if not trigger: return caption return f"{trigger}, {caption}" if caption else trigger