File size: 8,535 Bytes
9c98083
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
"""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