bjooo's picture
Upload folder using huggingface_hub
9c98083 verified
Raw History Blame Contribute Delete
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)
@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