hallucination / experiment /evaluation /eval_efuf_val.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
25.2 kB
"""
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
# 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()