Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
AtlasVision / code /eval_stage2.py
UncleanCode's picture
Upload via givemeanode export_data
2a2540a verified
Raw History Blame
9.96 kB
#!/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")