ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
7.81 kB
from __future__ import annotations
import warnings
import torch
from PIL import Image
from tqdm import tqdm
def generate_text(model, processor, image: Image.Image, prompt: str, device: str,
max_new_tokens: int = 300) -> str:
return generate_text_batch(model, processor, [image], prompt, device, max_new_tokens)[0]
def generate_text_batch(
model,
processor,
images: list,
prompt: str,
device: str,
max_new_tokens: int = 300,
) -> list[str]:
"""Generate captions for a batch of images with the same prompt.
Uses left-padding (required for batched generation) and restores the
tokenizer's original padding side afterwards.
"""
texts = [f"USER: <image>\n{prompt}\nASSISTANT:" for _ in images]
_orig_side = processor.tokenizer.padding_side
processor.tokenizer.padding_side = "left"
try:
inputs = processor(images=images, text=texts, return_tensors="pt", padding=True)
finally:
processor.tokenizer.padding_side = _orig_side
inputs = {k: v.to(device) for k, v in inputs.items()}
input_len = inputs["input_ids"].shape[1]
with torch.no_grad():
output_ids = model.generate(
**inputs,
max_new_tokens=max_new_tokens,
do_sample=False,
pad_token_id=processor.tokenizer.pad_token_id,
)
new_ids = output_ids[:, input_len:]
return processor.batch_decode(new_ids, skip_special_tokens=True)
def _build_vllm_prompt(processor, prompt: str) -> str:
if hasattr(processor, "apply_chat_template"):
messages = [{
"role": "user",
"content": [
{"type": "image", "image": "placeholder"},
{"type": "text", "text": prompt},
],
}]
try:
return processor.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
except Exception:
pass
return f"<image>\n{prompt}"
def collect_outputs_transformers(
model,
processor,
categories: dict,
prompts: list[str],
max_new_tokens: int,
device: str,
label: str = "model",
) -> dict:
outputs = {}
for cat_name, images in categories.items():
cat_outputs = []
for entry in tqdm(images, desc=f" {label}/{cat_name}", leave=False):
if "image" in entry:
image = entry["image"].convert("RGB")
else:
image = Image.open(entry["path"]).convert("RGB")
for prompt in prompts:
text = generate_text(model, processor, image, prompt, device, max_new_tokens)
cat_outputs.append({
"image_id": entry["image_id"],
"image_path": entry.get("path", entry["image_id"]),
"prompt": prompt,
"text": text,
})
outputs[cat_name] = cat_outputs
return outputs
def collect_outputs_visedit(
editor,
categories: dict,
prompts: list[str],
max_new_tokens: int,
edit_targets: dict | None = None,
label: str = "visedit",
relation_config=None,
) -> dict:
"""Collect outputs using VisEdit (VEAD) single-edit inference.
For the efficacy category (e.g. bathroom_no_toilet): applies edit_one_piece
(sets edit signal) before each generation, then restores.
For all other categories: plain inference.
Args:
editor: Loaded VEAD editor (from utils.load_vllm_editor).
categories: {cat_name: [{"image_id": ..., "image": PIL | "path": str}]}.
edit_targets: {image_id: target_new} for efficacy category images.
Used as the correction target when computing the edit signal.
Falls back to a generic description if not provided.
label: Display label for tqdm.
relation_config: RelationConfig for this relation.
"""
if relation_config is not None:
EDIT_CAT = relation_config.efficacy_category
DEFAULT_TARGET = (
f"A {relation_config.scene_key.replace('_', ' ')} scene "
f"without a {relation_config.object_key.replace('_', ' ')}."
)
else:
DEFAULT_TARGET = "A clean bathroom with a sink and mirror, without a toilet."
EDIT_CAT = "bathroom_no_toilet"
model = editor.vllm.model
processor = editor.vllm.processor
device = editor.device
outputs = {}
for cat_name, images in categories.items():
apply_edit = cat_name == EDIT_CAT
cat_outputs = []
for entry in tqdm(images, desc=f" {label}/{cat_name}", leave=False):
if "image" in entry:
image = entry["image"].convert("RGB")
else:
image = Image.open(entry["path"]).convert("RGB")
if apply_edit:
target = (edit_targets or {}).get(entry["image_id"], DEFAULT_TARGET)
request = {
"image": image,
"prompt": prompts[0],
"target_new": target,
}
editor.edit_one_piece(request)
for prompt in prompts:
text = generate_text(model, processor, image, prompt, device, max_new_tokens)
cat_outputs.append({
"image_id": entry["image_id"],
"image_path": entry.get("path", entry["image_id"]),
"prompt": prompt,
"text": text,
})
if apply_edit:
editor.restore_to_original_model()
outputs[cat_name] = cat_outputs
return outputs
def collect_outputs_vllm(
llm,
lora_request,
processor,
categories: dict,
prompts: list[str],
max_new_tokens: int,
batch_size: int,
label: str = "model",
) -> dict:
from vllm import SamplingParams
prompt_texts = {p: _build_vllm_prompt(processor, p) for p in prompts}
sampling_params = SamplingParams(max_tokens=max_new_tokens, temperature=0)
outputs = {}
for cat_name, images in categories.items():
cat_outputs = []
requests = [(entry, prompt) for entry in images for prompt in prompts]
for i in tqdm(range(0, len(requests), batch_size), desc=f" {label}/{cat_name}", leave=False):
batch = requests[i:i + batch_size]
vllm_inputs = []
contexts = []
for entry, prompt in batch:
try:
if "image" in entry:
image = entry["image"].convert("RGB")
else:
image = Image.open(entry["path"]).convert("RGB")
except Exception as exc:
warnings.warn(f"Image load failed for {entry.get('path', entry['image_id'])}: {exc}")
continue
vllm_inputs.append({
"prompt": prompt_texts[prompt],
"multi_modal_data": {"image": image},
})
contexts.append((entry, prompt))
if not vllm_inputs:
continue
generate_kwargs = {"sampling_params": sampling_params}
if lora_request is not None:
generate_kwargs["lora_request"] = lora_request
batch_outputs = llm.generate(vllm_inputs, **generate_kwargs)
for (entry, prompt), out in zip(contexts, batch_outputs):
text = out.outputs[0].text if out.outputs else ""
cat_outputs.append({
"image_id": entry["image_id"],
"image_path": entry.get("path", entry["image_id"]),
"prompt": prompt,
"text": text,
})
outputs[cat_name] = cat_outputs
return outputs