#!/usr/bin/env python3 """Build the de-biasing mix for stage 2b from LLaVA-1.5's public training mix (COCO-image parts only). Keeps VQAv2 / OK-VQA / A-OKVQA style short-answer conversations (many "no" answers) plus a replay of LLaVA-Instruct conversations, drops region/bounding-box tasks, and removes every POPE test image and the 1,000 held-out stage-2 conversations so evaluation stays clean.""" import collections import glob import json import os import random import re import pyarrow.parquet as pq import torch from huggingface_hub import hf_hub_download H = os.path.expanduser("~") D = f"{H}/data/llava_instruct" N_VQA = int(os.environ.get("N_VQA", 110_000)) N_LLAVA = int(os.environ.get("N_LLAVA", 40_000)) OUT = os.environ.get("MIX_OUT", f"{D}/mix_debias.json") mix_path = os.environ.get("MIX_PATH") or hf_hub_download("liuhaotian/LLaVA-Instruct-150K", "llava_v1_5_mix665k.json", repo_type="dataset", local_dir=D) mix = json.load(open(mix_path)) print(f"mix665k: {len(mix):,} entries") inst = json.load(open(f"{D}/llava_instruct_150k.json")) # same split as train_stage2.py (seed 42, last 1000) perm = torch.randperm(len(inst), generator=torch.Generator().manual_seed(42)).tolist() held_imgs = {inst[i]["image"] for i in perm[-int(os.environ.get("HELDOUT_N", 1000)):]} pope_ids = set() for f in glob.glob(f"{H}/data/pope/**/*.parquet", recursive=True): for src in pq.read_table(f, columns=["image_source"]).column("image_source").to_pylist(): m = re.search(r"(\d+)$", str(src)) if m: pope_ids.add(int(m.group(1))) print(f"POPE images to exclude: {len(pope_ids):,} | held-out stage-2 images to exclude: {len(held_imgs):,}") BBOX = re.compile(r"\[\s*\d?\.\d+\s*,\s*\d?\.\d+") SHORT = ("single word or phrase", "option's letter", "Answer the question using a single word") stats = collections.Counter() vqa, llava = [], [] for e in mix: img = e.get("image") or "" if not img.startswith("coco/train2017/"): stats["skip_not_coco"] += 1 continue base = img.rsplit("/", 1)[-1] if int(base.split(".")[0]) in pope_ids: stats["skip_pope_image"] += 1 continue if base in held_imgs: stats["skip_heldout_image"] += 1 continue text = " ".join(t["value"] for t in e["conversations"]) if BBOX.search(text): stats["skip_region_task"] += 1 continue item = {"id": str(e.get("id", base)), "image": base, "conversations": e["conversations"]} (vqa if any(s in text for s in SHORT) else llava).append(item) rng = random.Random(42) rng.shuffle(vqa) rng.shuffle(llava) out = vqa[:N_VQA] + llava[:N_LLAVA] rng.shuffle(out) json.dump(out, open(OUT, "w")) answers = collections.Counter() for e in vqa[:N_VQA]: for t in e["conversations"]: if t["from"] == "gpt": a = t["value"].strip().lower().rstrip(".") answers["yes" if a == "yes" else "no" if a == "no" else "other"] += 1 print("filter stats:", dict(stats)) print(f"available: {len(vqa):,} short-answer + {len(llava):,} instruct | using {min(N_VQA, len(vqa)):,} + {min(N_LLAVA, len(llava)):,} = {len(out):,}") print(f"short-answer turns: yes {answers['yes']:,} | no {answers['no']:,} | other {answers['other']:,}") if len(out) < 1000 and not os.environ.get("ALLOW_SMALL"): raise SystemExit(f"only {len(out)} examples - something is wrong with the filters") print(f"wrote {OUT}")