Download custom_nodes/dolphin_nodes/dataset_builder/caption.py from bjooo/tutorials: direct link, hf CLI and curl.
- Browser
- Download file 8.54 kB
-
https://huggingface.co/bjooo/tutorials/resolve/main/custom_nodes/dolphin_nodes/dataset_builder/caption.py
- Command line
-
hf download hf://bjooo/tutorials/custom_nodes/dolphin_nodes/dataset_builder/caption.py
-
curl -L -o caption.py https://huggingface.co/bjooo/tutorials/resolve/main/custom_nodes/dolphin_nodes/dataset_builder/caption.py
8.54 kB
| """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, <MORE_DETAILED_CAPTION> 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 = "<MORE_DETAILED_CAPTION>" | |
| 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) | |
| 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 | |