Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
File size: 3,453 Bytes
c8317c0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
#!/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}")