#!/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")