""" EFUF val-split eval: keyword mention rate — base vs EFUF MM-projector fine-tuned model. Loads LLaVA v1.5-7b (original liuhaotian format), optionally overlays an EFUF checkpoint (MM projector weights), generates captions on a relation's val split, and compares hallucination rates across 4 categories: - {scene}_no_{obj}: scene present, object absent (hallucination target) - {scene}_with_{obj}: scene present, object present (should still mention) - non_{scene}_with_{obj}: scene absent, object present (specificity) - neither: scene absent, object absent Usage: # Base model only (no EFUF checkpoint): python experiment/evaluation/eval_efuf_val.py # EFUF checkpoint: python experiment/evaluation/eval_efuf_val.py \ --efuf_ckpt path/to/step_007200.pth # Specifying prompts and GPU: python experiment/evaluation/eval_efuf_val.py \ --efuf_ckpt ... --prompts "Describe this image." --device cuda:2 """ from __future__ import annotations import argparse import json import os import sys import numpy as np import torch import torch.distributed as dist from tqdm import tqdm LLAVA_PATH = "/home/erwin/.cache/huggingface/hub/models--liuhaotian--llava-v1.5-7b/snapshots/4481d270cc22fd5c4d1bb5df129622006ccd9234" EFUF_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "..", "EFUF", "efuf") EXPERIMENT_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..") DEFAULT_EVAL_PROMPTS = [ "Describe this image.", "list all objects in this image", ] _CHAT_PREFIX = ( "A chat between a curious user and an artificial intelligence assistant. " "The assistant gives helpful, detailed, and polite answers to the user's questions. " ) def _configure_efuf_args(llava_path, device, max_new_tokens): """Inject EFUF-compatible args into sys.argv before importing common.args.""" efuf_argv = [ "--model", "llava", "--llava_path", llava_path, "--llava_ckpt_load_path", llava_path, "--device", device, "--max_new_tokens", str(max_new_tokens), "--llava_data_size_k", "0", "--gold_w", "0", "--sent_w", "0", "--run_name", "eval_only", ] saved = sys.argv[:] sys.argv = ["eval_efuf_val"] + efuf_argv return saved def _restore_argv(saved): sys.argv = saved def _efuf_generate_batch(llava_model_obj, model, vis_processor, images, prompt, device, max_new_tokens): """Generate captions for a batch of images with a single prompt (greedy, no beam search). Uses left-padding (required for batched generate) and passes attention_mask so the model ignores pad tokens during generation. """ texts = [f"{_CHAT_PREFIX}USER: \n{prompt} ASSISTANT:" for _ in images] ids_list = [llava_model_obj.tokenize_image(t) for t in texts] max_len = max(ids.shape[0] for ids in ids_list) pad_id = llava_model_obj.tokenizer.pad_token_id padded_ids, attn_masks = [], [] for ids in ids_list: pad_len = max_len - ids.shape[0] if pad_len > 0: padding = torch.full((pad_len,), pad_id, dtype=ids.dtype) padded_ids.append(torch.cat([padding, ids])) attn_masks.append(torch.cat([torch.zeros(pad_len, dtype=torch.long), torch.ones(ids.shape[0], dtype=torch.long)])) else: padded_ids.append(ids) attn_masks.append(torch.ones(ids.shape[0], dtype=torch.long)) input_ids = torch.stack(padded_ids).to(device) attention_mask = torch.stack(attn_masks).to(device) pixel_values = torch.stack([vis_processor(img) for img in images]).to(device, model.dtype) with torch.inference_mode(): output_ids = model.generate( input_ids=input_ids, images=pixel_values, attention_mask=attention_mask, do_sample=False, pad_token_id=pad_id, max_new_tokens=max_new_tokens, ) new_ids = output_ids[:, input_ids.shape[1]:] return [c.strip() for c in llava_model_obj.tokenizer.batch_decode(new_ids, skip_special_tokens=True)] def parse_args(): p = argparse.ArgumentParser(description="Evaluate EFUF fine-tuned LLaVA on hallucination metrics") p.add_argument("--relation", type=str, default="bathroom_toilet") p.add_argument("--llava_path", type=str, default=LLAVA_PATH) p.add_argument("--efuf_ckpt", type=str, default="", help="Path to EFUF checkpoint .pth. Empty = base model only.") p.add_argument("--output_dir", type=str, default=None) p.add_argument("--prompts", type=str, nargs="+", default=None) p.add_argument("--max_new_tokens", type=int, default=300) p.add_argument("--max_samples", type=int, default=0, help="Max samples per category. 0 = full val split.") p.add_argument("--device", type=str, default="cuda:0") p.add_argument("--dtype", type=str, default="float16", choices=["float16", "bfloat16"]) p.add_argument("--skip_base", action="store_true", help="Skip base-model inference; only run EFUF model.") p.add_argument("--split", type=str, default="validation", help="HuggingFace dataset split.") p.add_argument("--seed", type=int, default=42) p.add_argument("--batch_size", type=int, default=1, help="Images per generate call.") p.add_argument("--num_shards", type=int, default=1, help="Total shards for parallel inference.") p.add_argument("--shard_rank", type=int, default=0, help="This shard's rank (0-indexed).") return p.parse_args() def _four_category_masks(sc_arr, ob_arr): return ( (sc_arr < 0.5) & (ob_arr > 0.5), (sc_arr > 0.5) & (ob_arr < 0.5), (sc_arr > 0.5) & (ob_arr > 0.5), (sc_arr < 0.5) & (ob_arr < 0.5), ) def _bleu_per_category(base_caps: list, edited_caps: list, sc_arr, ho_arr, cat_order: list) -> dict: try: from nltk.translate.bleu_score import sentence_bleu, SmoothingFunction _smooth = SmoothingFunction().method1 def _score(ref: str, hyp: str) -> float: r, h = ref.lower().split(), hyp.lower().split() if not r or not h: return float("nan") return sentence_bleu([r], h, weights=(0.5, 0.5), smoothing_function=_smooth) except ImportError: def _score(ref: str, hyp: str) -> float: r, h = set(ref.lower().split()), set(hyp.lower().split()) if not r or not h: return float("nan") inter = len(r & h) p, rec = inter / len(h), inter / len(r) return 2 * p * rec / (p + rec) if (p + rec) > 0 else 0.0 m_to, m_bo, m_bt, m_ne = _four_category_masks(sc_arr, ho_arr) masks = dict(zip(cat_order, (m_to, m_bo, m_bt, m_ne))) out = {} for cat, mask in masks.items(): if not mask.any(): out[cat] = None continue scores = [_score(base_caps[i], edited_caps[i]) for i in np.where(mask)[0]] valid = [s for s in scores if not np.isnan(s)] out[cat] = float(np.mean(valid)) if valid else None return out def _sample_cat(sc_v, ob_v, cat_obj_only, cat_scene_only, cat_both, cat_neither): if sc_v < 0.5 and ob_v > 0.5: return cat_obj_only if sc_v > 0.5 and ob_v < 0.5: return cat_scene_only if sc_v > 0.5 and ob_v > 0.5: return cat_both return cat_neither def _setup_dist(): """Initialize distributed if torchrun set LOCAL_RANK; otherwise single-process.""" if "LOCAL_RANK" not in os.environ: return 0, 1, None local_rank = int(os.environ["LOCAL_RANK"]) torch.cuda.set_device(local_rank) dist.init_process_group(backend="nccl") return dist.get_rank(), dist.get_world_size(), local_rank def _gather_list(local: list, world_size: int) -> list: if world_size == 1: return local bucket = [None] * world_size dist.all_gather_object(bucket, local) out = [] for part in bucket: out.extend(part) return out def _json_default(obj): if isinstance(obj, np.generic): return obj.item() if isinstance(obj, float) and np.isnan(obj): return None raise TypeError(f"Object of type {type(obj)} is not JSON serializable") def main(): parsed = parse_args() prompts = parsed.prompts if parsed.prompts else list(DEFAULT_EVAL_PROMPTS) batch_size = parsed.batch_size rank, world_size, local_rank_ddp = _setup_dist() is_main = rank == 0 # DDP overrides --device; fall back to manual --num_shards/--shard_rank otherwise if local_rank_ddp is not None: device = f"cuda:{local_rank_ddp}" else: device = parsed.device sys.path.insert(0, EXPERIMENT_DIR) from config.relation_config import get_relation_config from data.hf_loader import load_hf_dataset from evaluation.metrics import KeywordMentionDetector rc = get_relation_config(parsed.relation) scene_col, obj_col = rc.scene_key, rc.object_key cat_obj_only = rc.non_scene_with_object cat_scene_only = rc.scene_no_object cat_both = rc.scene_with_object cat_neither = "neither" cat_order = [cat_obj_only, cat_scene_only, cat_both, cat_neither] cat_display = { rc.non_scene_with_object: f"{rc.object_key}_only", rc.scene_no_object: f"{rc.scene_key}_no_{rc.object_key}", rc.scene_with_object: f"{rc.scene_key}_with_{rc.object_key}", "neither": "neither", } kw = KeywordMentionDetector(keywords=rc.mention_keywords) ds = load_hf_dataset(rc.dataset_id, split=parsed.split) if is_main: print(f"Loaded {rc.dataset_id} split={parsed.split}: {len(ds)} samples") sc_labels = ds[scene_col] ob_labels = ds[obj_col] if parsed.max_samples <= 0: all_indices = list(range(len(ds))) else: rng = np.random.default_rng(parsed.seed) cat_buckets = {cat_obj_only: [], cat_scene_only: [], cat_both: [], cat_neither: []} for i, (sc_v, ob_v) in enumerate(zip(sc_labels, ob_labels)): if int(sc_v) == 0 and int(ob_v) == 1: cat_buckets[cat_obj_only].append(i) elif int(sc_v) == 1 and int(ob_v) == 0: cat_buckets[cat_scene_only].append(i) elif int(sc_v) == 1 and int(ob_v) == 1: cat_buckets[cat_both].append(i) else: cat_buckets[cat_neither].append(i) all_indices = [] for bucket in cat_buckets.values(): all_indices.extend(bucket[:parsed.max_samples]) all_indices.sort() # DDP: stripe indices across ranks. Legacy: chunk via --num_shards/--shard_rank. if world_size > 1: indices = [all_indices[i] for i in range(rank, len(all_indices), world_size)] if is_main: print(f"DDP world_size={world_size}, each rank processes ~{len(indices)} samples (striped)") elif parsed.num_shards > 1: shard_size = (len(all_indices) + parsed.num_shards - 1) // parsed.num_shards start = parsed.shard_rank * shard_size indices = all_indices[start:start + shard_size] print(f"Shard {parsed.shard_rank}/{parsed.num_shards}: indices [{start}:{start + len(indices)}]") else: indices = all_indices n_total = len(all_indices) n = len(indices) if is_main: print(f"Total eval samples: {n_total} this rank: {n} batch_size: {batch_size}") saved_argv = _configure_efuf_args(parsed.llava_path, device, parsed.max_new_tokens) sys.path.insert(0, EFUF_DIR) from common.args import args as efuf_args from common.models import LlavaModel, load_ckpt _restore_argv(saved_argv) if is_main: print(f"Loading base LLaVA model from {parsed.llava_path} ...") llava_model_obj = LlavaModel() model, vis_processor = llava_model_obj.load(parsed.llava_path, str(device), train=False) model.eval() base_captions: list[str] = [] base_rates: list[float] = [] if not parsed.skip_base: if is_main: print("\n=== (1) Base model caption generation ===") for _bs in tqdm(range(0, n, batch_size), desc="Base model", unit="batch", dynamic_ncols=True, disable=not is_main): batch_idx = indices[_bs:_bs + batch_size] batch_rows = [ds[i] for i in batch_idx] batch_images = [r["image"].convert("RGB") for r in batch_rows] for prompt in prompts: captions = _efuf_generate_batch( llava_model_obj, model, vis_processor, batch_images, prompt, device, parsed.max_new_tokens ) for caption in captions: base_captions.append(caption) base_rates.append(float(kw.mentions_object(caption))) if parsed.efuf_ckpt: if is_main: print(f"\n=== Loading EFUF checkpoint: {parsed.efuf_ckpt} ===") checkpoint = torch.load(parsed.efuf_ckpt, map_location=str(device)) state_dict = checkpoint["model"] if "model" in checkpoint else checkpoint model.load_state_dict(state_dict, strict=False) model.eval() efuf_label = "EFUF" if parsed.efuf_ckpt else "base" efuf_captions: list[str] = [] efuf_rates: list[float] = [] gt_has_object: list[float] = [] scene_flags: list[int] = [] image_ids: list = [] eval_indices: list[int] = [] eval_prompts: list[str] = [] if is_main: print(f"\n=== (2) {efuf_label} model caption generation ===") for _bs in tqdm(range(0, n, batch_size), desc=efuf_label, unit="batch", dynamic_ncols=True, disable=not is_main): batch_idx = indices[_bs:_bs + batch_size] batch_rows = [ds[i] for i in batch_idx] batch_images = [r["image"].convert("RGB") for r in batch_rows] for prompt in prompts: captions = _efuf_generate_batch( llava_model_obj, model, vis_processor, batch_images, prompt, device, parsed.max_new_tokens ) for caption, idx, row in zip(captions, batch_idx, batch_rows): efuf_captions.append(caption) efuf_rates.append(float(kw.mentions_object(caption))) gt_has_object.append(float(int(row[obj_col]))) scene_flags.append(int(row[scene_col])) image_ids.append(row.get("image_id") if hasattr(row, "get") else None) eval_indices.append(idx) eval_prompts.append(prompt) # --- Gather across DDP ranks --- efuf_captions = _gather_list(efuf_captions, world_size) efuf_rates = _gather_list(efuf_rates, world_size) gt_has_object = _gather_list(gt_has_object, world_size) scene_flags = _gather_list(scene_flags, world_size) image_ids = _gather_list(image_ids, world_size) eval_indices = _gather_list(eval_indices, world_size) eval_prompts = _gather_list(eval_prompts, world_size) base_captions = _gather_list(base_captions, world_size) base_rates = _gather_list(base_rates, world_size) if not is_main: if dist.is_initialized(): dist.barrier() dist.destroy_process_group() return # --- Rank 0: compute metrics and write outputs --- l_arr = np.array(efuf_rates, dtype=np.float64) b_arr = np.array(base_rates, dtype=np.float64) if base_rates else None ho_arr = np.array(gt_has_object, dtype=np.float64) sc_arr = np.array(scene_flags, dtype=np.float64) # Sort by (index, prompt position) to get deterministic order _prompt_pos = {p: i for i, p in enumerate(prompts)} sort_order = sorted( range(len(eval_indices)), key=lambda j: (eval_indices[j], _prompt_pos.get(eval_prompts[j], 0)), ) efuf_captions = [efuf_captions[j] for j in sort_order] efuf_rates = [efuf_rates[j] for j in sort_order] gt_has_object = [gt_has_object[j] for j in sort_order] scene_flags = [scene_flags[j] for j in sort_order] image_ids = [image_ids[j] for j in sort_order] eval_indices = [eval_indices[j] for j in sort_order] eval_prompts = [eval_prompts[j] for j in sort_order] l_arr = np.array(efuf_rates, dtype=np.float64) ho_arr = np.array(gt_has_object, dtype=np.float64) sc_arr = np.array(scene_flags, dtype=np.float64) if base_rates: base_rates = [base_rates[j] for j in sort_order] base_captions = [base_captions[j] for j in sort_order] b_arr = np.array(base_rates, dtype=np.float64) m_to, m_bo, m_bt, m_ne = _four_category_masks(sc_arr, ho_arr) masks = dict(zip(cat_order, (m_to, m_bo, m_bt, m_ne))) _has_base = b_arr is not None n_images_total = len(set(eval_indices)) print(f"\n images: {n_images_total} evals: {len(l_arr)} prompts: {prompts!r}") print(f" keywords: {rc.mention_keywords[:3]!r}... (negation-aware)") if not _has_base: print(" (base-model columns omitted)") print() cap_metrics = {} cap_overall = {} _sep = "-" * (91 if _has_base else 60) _hdr = f"{'Base':>14} " if _has_base else "" print(f" {'Category':<24} {'N':>5} {_hdr}{efuf_label:>14} Error") print(" " + _sep) for cat in cat_order: mask = masks[cat] if not mask.any(): display_name = cat_display.get(cat, cat) print(f" {display_name:<24} {'(empty)'}") continue display_name = cat_display.get(cat, cat) nr = int(mask.sum()) lr = float(l_arr[mask].mean()) lm = int(l_arr[mask].sum()) has_obj = bool((ho_arr[mask] > 0.5).all()) error_type = "miss_rate" if has_obj else "hallu_rate" efuf_err = (1.0 - lr) if has_obj else lr _base_col = "" _be_s = "" if _has_base: br = float(b_arr[mask].mean()) bm = int(b_arr[mask].sum()) base_err = (1.0 - br) if has_obj else br _base_col = f"{bm:>3}/{nr:<4}({br:>6.1%}) " _be_s = f"base={base_err:.1%} " print( f" {display_name:<24} {nr:>5} " f"{_base_col}" f"{lm:>3}/{nr:<4}({lr:>6.1%}) " f"{error_type}: {_be_s}{efuf_label}={efuf_err:.1%}" ) cap_metrics[cat] = { "n": nr, "mention_count": lm, "mention_rate": lr, f"{efuf_label}_{error_type}": efuf_err, } if _has_base: cap_metrics[cat].update({ "base_mention_count": bm, "base_mention_rate": br, f"base_{error_type}": base_err, }) print(" " + _sep) overall_lr = float(l_arr.mean()) _overall_base = "" if _has_base: overall_br = float(b_arr.mean()) _overall_base = f"{int(b_arr.sum()):>3}/{len(b_arr):<4}({overall_br:>6.1%}) " print( f" {'OVERALL':<24} {len(l_arr):>5} " f"{_overall_base}" f"{int(l_arr.sum()):>3}/{len(l_arr):<4}({overall_lr:>6.1%})" ) cap_overall = {f"{efuf_label}_mention_rate": overall_lr} if _has_base: cap_overall["base_mention_rate"] = overall_br if masks[cat_scene_only] is not None and masks[cat_scene_only].any(): sd = cap_metrics.get(cat_scene_only, {}) hallu_key = f"{efuf_label}_hallu_rate" if hallu_key in sd: _sup_base = f"base hallu={sd.get('base_hallu_rate', float('nan')):.1%} " if _has_base else "" _delta = f" Δ={sd['base_hallu_rate'] - sd[hallu_key]:+.1%}" if _has_base else "" print( f"\n [Suppression] {cat_display.get(cat_scene_only, cat_scene_only)} " f"(D_{{scene,¬obj}}): {_sup_base}" f"{efuf_label} hallu={sd[hallu_key]:.1%}" f"{_delta}" ) print("\n Per-prompt overall:") prompt_arr = np.array(eval_prompts, dtype=object) cap_metrics_by_prompt = {} for prompt in prompts: pmask = prompt_arr == prompt p_l = l_arr[pmask] p_lr = float(p_l.mean()) if _has_base: p_br = float(b_arr[pmask].mean()) print(f" {prompt!r}: Base={p_br:.1%} {efuf_label}={p_lr:.1%}") else: print(f" {prompt!r}: {efuf_label}={p_lr:.1%}") pr_metrics = {"overall": {f"{efuf_label}_mention_rate": p_lr}} if _has_base: pr_metrics["overall"]["base_mention_rate"] = float(b_arr[pmask].mean()) cap_metrics_by_prompt[prompt] = pr_metrics bleu_vs_base: dict = {} if _has_base: bleu_vs_base = _bleu_per_category(base_captions, efuf_captions, sc_arr, ho_arr, cat_order) print("\n Caption similarity (edited vs base, BLEU-2) by category:") for cat in cat_order: v = bleu_vs_base.get(cat) if v is not None: print(f" {cat_display.get(cat, cat):<24} {v:.4f}") out_dir = parsed.output_dir if out_dir is None: base_dir = os.path.dirname(parsed.efuf_ckpt) if parsed.efuf_ckpt else "efuf_eval_results" out_dir = base_dir os.makedirs(out_dir, exist_ok=True) captions_records = [] for j in range(len(eval_indices)): rec = { "index": int(eval_indices[j]), "image_id": image_ids[j], "prompt": eval_prompts[j], scene_col: int(sc_arr[j]), obj_col: int(ho_arr[j]), "category": cat_display.get( _sample_cat(sc_arr[j], ho_arr[j], cat_obj_only, cat_scene_only, cat_both, cat_neither), _sample_cat(sc_arr[j], ho_arr[j], cat_obj_only, cat_scene_only, cat_both, cat_neither), ), f"{efuf_label}_caption": efuf_captions[j], f"{efuf_label}_mentions_object": bool(l_arr[j] > 0.5), } if _has_base: rec["base_caption"] = base_captions[j] rec["base_mentions_object"] = bool(b_arr[j] > 0.5) captions_records.append(rec) captions_path = os.path.join(out_dir, "captions.json") with open(captions_path, "w") as f: json.dump(captions_records, f, indent=2) print(f"\n Captions saved to {captions_path}") metrics = { "relation": parsed.relation, "efuf_ckpt": parsed.efuf_ckpt, "has_base": _has_base, "n_images": n_images_total, "n_prompt_evals": len(l_arr), "prompts": prompts, "caption_eval": { "overall": cap_overall, "bleu_vs_base": bleu_vs_base if bleu_vs_base else None, "categories": {cat: cap_metrics.get(cat) for cat in cat_order}, "per_prompt": cap_metrics_by_prompt, }, } metrics_path = os.path.join(out_dir, "metrics.json") with open(metrics_path, "w") as f: json.dump(metrics, f, indent=2, default=_json_default) print(f" Metrics saved to {metrics_path}") samples_by_idx = {} for j in range(len(eval_indices)): idx = int(eval_indices[j]) sc_v = float(sc_arr[j]) ob_v = float(ho_arr[j]) cat = _sample_cat(sc_v, ob_v, cat_obj_only, cat_scene_only, cat_both, cat_neither) rec = samples_by_idx.setdefault(idx, { "index": idx, "image_id": image_ids[j], scene_col: int(sc_v), obj_col: int(ob_v), "category": cat_display.get(cat, cat), "prompt_results": {}, }) prompt_rec = rec["prompt_results"].setdefault(eval_prompts[j], {}) prompt_rec[f"{efuf_label}_caption"] = efuf_captions[j] prompt_rec[f"{efuf_label}_mentions_object"] = bool(l_arr[j] > 0.5) if _has_base: prompt_rec["base_caption"] = base_captions[j] prompt_rec["base_mentions_object"] = bool(b_arr[j] > 0.5) if len(prompts) == 1: rec[f"{efuf_label}_caption"] = efuf_captions[j] rec[f"{efuf_label}_mentions_object"] = bool(l_arr[j] > 0.5) if _has_base: rec["base_caption"] = base_captions[j] rec["base_mentions_object"] = bool(b_arr[j] > 0.5) samples_path = os.path.join(out_dir, "samples.json") with open(samples_path, "w") as f: json.dump([samples_by_idx[k] for k in sorted(samples_by_idx)], f, indent=2, default=_json_default) print(f" Samples saved to {samples_path}") # Legacy: write shard_meta.json when called with --num_shards so merge_shards.py still works if parsed.num_shards > 1 and world_size == 1: meta = { "relation": parsed.relation, "efuf_ckpt": parsed.efuf_ckpt, "efuf_label": efuf_label, "has_base": _has_base, "prompts": prompts, "scene_col": scene_col, "obj_col": obj_col, "cat_obj_only": cat_obj_only, "cat_scene_only": cat_scene_only, "cat_both": cat_both, "cat_neither": cat_neither, "cat_order": cat_order, "cat_display": cat_display, "num_shards": parsed.num_shards, "shard_rank": parsed.shard_rank, "n_shard": n, } with open(os.path.join(out_dir, "shard_meta.json"), "w") as f: json.dump(meta, f, indent=2) print(f" Shard {parsed.shard_rank} done. Meta saved.") if dist.is_initialized(): dist.barrier() dist.destroy_process_group() if __name__ == "__main__": main()