Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
File size: 9,964 Bytes
2a2540a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
#!/usr/bin/env python3
"""Evaluate stage 1 vs stage 2 and write ~/eval/stage2_eval.json + ~/eval/report.md
1) held-out LLaVA-Instruct loss   2) POPE (random / popular / adversarial): acc, precision, recall, F1, yes-ratio
3) detailed descriptions on POPE images   4) image questions in Igbo / Yoruba / Hausa   5) text-only check of N-ATLaS
"""
import glob
import io
import json
import os
import time

os.environ.setdefault("HF_HUB_OFFLINE", "1")
import torch
from PIL import Image
from torch.utils.data import DataLoader, Subset

import train as T
import train_stage2 as S

HOME = os.path.expanduser("~")
OUT = T.env("EVAL_DIR", f"{HOME}/eval")
POPE_DIR = T.env("POPE_DIR", f"{HOME}/data/pope")
POPE_LIMIT = T.env("POPE_LIMIT", 0, int)  # 0 = all questions per split
HELD_N = T.env("HELD_N", 500, int)
os.makedirs(OUT, exist_ok=True)
DEV, DT = T.DEVICE, T.DTYPE
t0 = time.time()

proj1 = S._load_proj(S.STAGE1_PROJECTOR)
proj2 = torch.load(os.path.join(S.CKPT_DIR, "projector_stage2.pt"), map_location="cpu")
from peft import PeftModel

model, tok, pad_id, processor = T.build_model()
model.llm = PeftModel.from_pretrained(model.llm, os.path.join(S.CKPT_DIR, "lora_adapter")).to(DEV)
model.llm.eval()
model.eval()
EOT = tok.convert_tokens_to_ids(T.EOT)


class Stage:
    """Context manager: stage 1 = stage-1 projector, adapters off; stage 2 = stage-2 projector, adapters on."""
    def __init__(self, n):
        self.n = n

    def __enter__(self):
        model.projector.load_state_dict(proj1 if self.n == 1 else proj2)
        model.projector.to(DEV, dtype=torch.float32)
        self.ctx = model.llm.disable_adapter() if self.n == 1 else None
        if self.ctx:
            self.ctx.__enter__()

    def __exit__(self, *a):
        if self.ctx:
            self.ctx.__exit__(*a)


res = {"stage1": {}, "stage2": {}}

# ---------------- 1. held-out instruct loss ----------------
ann = json.load(open(S.INSTRUCT_JSON))
_, held = S.split(ann)
prefix = T.find_zip_prefix([ann[i] for i in held[:200]], S.COCO_ZIP)
ds = S.InstructDataset(ann, S.COCO_ZIP, prefix, processor, tok, S.MAX_TEXT_LEN)
dl = DataLoader(Subset(ds, held[:HELD_N]), batch_size=8, num_workers=T.NUM_WORKERS,
                collate_fn=T.make_collate(pad_id))


@torch.no_grad()
def heldout_loss():
    tot, n = 0.0, 0
    for b in dl:
        with torch.autocast(device_type=DEV.type, dtype=DT):
            l = model(b["pixel_values"].to(DEV), b["input_ids"].to(DEV), b["attention_mask"].to(DEV),
                      b["labels"].to(DEV)).float()
        k = int((b["labels"] != -100).sum())
        tot, n = tot + l.item() * k, n + k
    return round(tot / n, 4)


for s in (1, 2):
    with Stage(s):
        res[f"stage{s}"]["heldout_instruct_loss"] = heldout_loss()
T.log(f"held-out instruct loss: stage1 {res['stage1']['heldout_instruct_loss']} | "
      f"stage2 {res['stage2']['heldout_instruct_loss']}")

# ---------------- 2. POPE ----------------
import pyarrow.parquet as pq

SUFFIX = " Answer the question using a single word or phrase."
yes_ids = sorted({tok(w, add_special_tokens=False).input_ids[0] for w in ("Yes", "yes", " Yes", " yes")})
no_ids = sorted({tok(w, add_special_tokens=False).input_ids[0] for w in ("No", "no", " No", " no")})


def to_image(cell):
    if isinstance(cell, dict):
        cell = cell.get("bytes") or open(cell["path"], "rb").read()
    return Image.open(io.BytesIO(cell)).convert("RGB")


def pope_rows(split):
    files = sorted(glob.glob(f"{POPE_DIR}/**/{split}-*.parquet", recursive=True))
    if files:
        rows = pq.read_table(files[0]).to_pylist()
    else:  # some versions ship one 'test' table with a 'category' column
        rows = [r for f in sorted(glob.glob(f"{POPE_DIR}/**/test-*.parquet", recursive=True))
                for r in pq.read_table(f).to_pylist() if r.get("category") == split]
    return rows[:POPE_LIMIT] if POPE_LIMIT else rows


@torch.no_grad()
def pope_eval(rows, bs=32):
    tp = fp = tn = fn = 0
    for s in range(0, len(rows), bs):
        chunk = rows[s:s + bs]
        pv = torch.stack([processor(images=to_image(r["image"]), return_tensors="pt").pixel_values[0] for r in chunk]).to(DEV)
        seqs = [tok(r["question"].strip() + SUFFIX + T.ASSIST_HEADER, add_special_tokens=False).input_ids for r in chunk]
        L = max(map(len, seqs))
        ids = torch.full((len(seqs), L), pad_id, dtype=torch.long)
        mask = torch.zeros_like(ids)
        for i, q in enumerate(seqs):
            ids[i, :len(q)] = torch.tensor(q)
            mask[i, :len(q)] = 1
        ids, mask = ids.to(DEV), mask.to(DEV)
        with torch.autocast(device_type=DEV.type, dtype=DT):
            e, m, _ = model.build_inputs(pv, ids, mask)
            logits = model.llm(inputs_embeds=e, attention_mask=m).logits
        n_fixed = e.shape[1] - L
        for i, r in enumerate(chunk):
            last = logits[i, n_fixed + len(seqs[i]) - 1].float()
            pred_yes = last[yes_ids].max() > last[no_ids].max()
            gold_yes = str(r["answer"]).strip().lower().startswith("yes")
            tp += pred_yes and gold_yes
            fp += pred_yes and not gold_yes
            tn += (not pred_yes) and (not gold_yes)
            fn += (not pred_yes) and gold_yes
    tp, fp, tn, fn = map(int, (tp, fp, tn, fn))
    n = tp + fp + tn + fn
    prec = tp / max(1, tp + fp)
    rec = tp / max(1, tp + fn)
    return {"n": n, "accuracy": round((tp + tn) / max(1, n), 4), "precision": round(prec, 4),
            "recall": round(rec, 4), "f1": round(2 * prec * rec / max(1e-9, prec + rec), 4),
            "yes_ratio": round((tp + fp) / max(1, n), 4)}


pope_samples = []
for split in ("random", "popular", "adversarial"):
    rows = pope_rows(split)
    if not rows:
        T.log(f"POPE {split}: no data found")
        continue
    pope_samples = pope_samples or rows
    for s in (1, 2):
        with Stage(s):
            res[f"stage{s}"][f"pope_{split}"] = pope_eval(rows)
    T.log(f"POPE {split}: stage1 {res['stage1'][f'pope_{split}']} | stage2 {res['stage2'][f'pope_{split}']}")


# ---------------- 3/4. generations ----------------
@torch.no_grad()
def answer(image, question, max_new_tokens=120):
    pv = processor(images=image, return_tensors="pt").pixel_values.to(DEV)
    ids = torch.tensor([tok(question + T.ASSIST_HEADER, add_special_tokens=False).input_ids], device=DEV)
    with torch.autocast(device_type=DEV.type, dtype=DT):
        e, m, _ = model.build_inputs(pv, ids, torch.ones_like(ids))
        out = model.llm.generate(inputs_embeds=e, attention_mask=m, max_new_tokens=max_new_tokens,
                                 do_sample=False, eos_token_id=EOT, pad_token_id=pad_id, repetition_penalty=1.1)
    return tok.decode(out[0], skip_special_tokens=True).strip()


seen, gen_images = set(), []
for r in pope_samples:
    key = r.get("image_source") or r.get("question_id")
    if key not in seen:
        seen.add(key)
        gen_images.append((str(key), to_image(r["image"])))
    if len(gen_images) == 4:
        break

res["descriptions"] = []
for key, img in gen_images:
    row = {"image": key}
    for s in (1, 2):
        with Stage(s):
            row[f"stage{s}"] = answer(img, "Describe this image in detail.")
    res["descriptions"].append(row)

MULTI = {"igbo": "Kedu ihe dị na foto a?", "yoruba": "Kí ni ó wà nínú àwòrán yìí?", "hausa": "Me ke cikin wannan hoton?"}
res["multilingual"] = []
if gen_images:
    with Stage(2):
        for lang, q in MULTI.items():
            res["multilingual"].append({"image": gen_images[0][0], "language": lang, "question": q,
                                        "stage2": answer(gen_images[0][1], q)})

# ---------------- 5. text-only check ----------------
@torch.no_grad()
def text_only(q, max_new_tokens=120):
    ids = tok(T.USER_HEADER + q + T.ASSIST_HEADER, add_special_tokens=True, return_tensors="pt").input_ids.to(DEV)
    with torch.autocast(device_type=DEV.type, dtype=DT):
        out = model.llm.generate(input_ids=ids, attention_mask=torch.ones_like(ids), max_new_tokens=max_new_tokens,
                                 do_sample=False, eos_token_id=EOT, pad_token_id=pad_id, repetition_penalty=1.1)
    return tok.decode(out[0, ids.shape[1]:], skip_special_tokens=True).strip()


q_text = "Kedu ihe bụ positron? Kọwaa ya n'asụsụ Igbo."
with Stage(1):
    base_answer = text_only(q_text)
with Stage(2):
    lora_answer = text_only(q_text)
res["text_only"] = {"question": q_text, "base_n_atlas": base_answer, "with_stage2_lora": lora_answer}
res["eval_minutes"] = round((time.time() - t0) / 60, 1)
json.dump(res, open(os.path.join(OUT, "stage2_eval.json"), "w"), indent=1, ensure_ascii=False)

# ---------------- report ----------------
L = ["# Atlas-Vision evaluation\n", "## Scores (stage 1 → stage 2)\n", "| Metric | Stage 1 | Stage 2 |", "|---|---|---|",
     f"| Held-out instruct loss (lower is better) | {res['stage1']['heldout_instruct_loss']} | {res['stage2']['heldout_instruct_loss']} |"]
for split in ("random", "popular", "adversarial"):
    if f"pope_{split}" in res["stage1"]:
        a, b = res["stage1"][f"pope_{split}"], res["stage2"][f"pope_{split}"]
        L.append(f"| POPE {split}: accuracy / F1 / yes-ratio | {a['accuracy']} / {a['f1']} / {a['yes_ratio']} | "
                 f"{b['accuracy']} / {b['f1']} / {b['yes_ratio']} |")
L.append("\n## Detailed descriptions\n")
for d in res["descriptions"]:
    L += [f"**{d['image']}**\n", f"- Stage 1: {d['stage1']}", f"- Stage 2: {d['stage2']}\n"]
L.append("## Questions in Nigerian languages (stage 2)\n")
for m in res["multilingual"]:
    L.append(f"- **{m['language']}** — {m['question']}\n  → {m['stage2']}")
L += ["\n## Text-only check (no image)\n", f"Question: {q_text}\n",
      f"- Base N-ATLaS: {base_answer}", f"- With stage-2 LoRA: {lora_answer}"]
open(os.path.join(OUT, "report.md"), "w").write("\n".join(L) + "\n")
T.log(f"EVAL_DONE in {res['eval_minutes']} min -> {OUT}/stage2_eval.json, {OUT}/report.md")