| """ |
| 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: <image>\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 |
|
|
| |
| 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() |
|
|
| |
| 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) |
|
|
| |
| 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 |
|
|
| |
| 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) |
|
|
| |
| _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}") |
|
|
| |
| 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() |
|
|