summaryrefslogtreecommitdiff
path: root/worldalign/evaluate.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/evaluate.py')
-rw-r--r--worldalign/evaluate.py149
1 files changed, 149 insertions, 0 deletions
diff --git a/worldalign/evaluate.py b/worldalign/evaluate.py
new file mode 100644
index 0000000..6f97873
--- /dev/null
+++ b/worldalign/evaluate.py
@@ -0,0 +1,149 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+import re
+from collections import Counter
+
+import torch
+from transformers import AutoModelForCausalLM, AutoTokenizer
+
+from .common import (
+ dtype_for_device,
+ read_json,
+ retrieval_metrics,
+ write_json,
+)
+from .io import load_feature_pair, select_rows
+from .models import load_bridge, load_prefix
+
+
+TOKEN_RE = re.compile(r"[a-z0-9]+")
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--vision", default="artifacts/vision.pt")
+ p.add_argument("--text", default="artifacts/text.pt")
+ p.add_argument("--bridge", required=True)
+ p.add_argument("--prefix")
+ p.add_argument("--split", choices=["val", "test"], default="test")
+ p.add_argument("--device", default="cuda:1")
+ p.add_argument("--generation-samples", type=int, default=100)
+ p.add_argument("--max-new-tokens", type=int, default=32)
+ p.add_argument(
+ "--shuffle-mapped",
+ action="store_true",
+ help="Permute image-conditioned latents before retrieval/generation as a null control.",
+ )
+ p.add_argument("--seed", type=int, default=20260728)
+ p.add_argument("--output", default="artifacts/evaluation.json")
+ return p.parse_args()
+
+
+def unigram_f1(candidate: str, references: list[str]) -> float:
+ candidate_tokens = TOKEN_RE.findall(candidate.lower())
+ if not candidate_tokens:
+ return 0.0
+ candidate_count = Counter(candidate_tokens)
+ best = 0.0
+ for reference in references:
+ reference_count = Counter(TOKEN_RE.findall(reference.lower()))
+ overlap = sum((candidate_count & reference_count).values())
+ precision = overlap / max(sum(candidate_count.values()), 1)
+ recall = overlap / max(sum(reference_count.values()), 1)
+ f1 = 2 * precision * recall / max(precision + recall, 1e-12)
+ best = max(best, f1)
+ return best
+
+
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(args.manifest)
+ vision, text, vlookup, tlookup = load_feature_pair(args.vision, args.text)
+ rows = manifest[args.split]
+ x = select_rows(vision["features"], vlookup, rows)
+ y = select_rows(text["features"], tlookup, rows)
+
+ bridge, bridge_state = load_bridge(args.bridge, args.device)
+ mapped = []
+ with torch.inference_mode():
+ for chunk in x.split(512):
+ mapped.append(bridge(chunk.to(args.device)).cpu())
+ mapped = torch.cat(mapped)
+ if args.shuffle_mapped:
+ generator = torch.Generator().manual_seed(args.seed)
+ mapped = mapped[torch.randperm(len(mapped), generator=generator)]
+ result: dict = {
+ "split": args.split,
+ "samples": len(rows),
+ "bridge_mode": bridge_state["mode"],
+ "shuffle_mapped": args.shuffle_mapped,
+ "retrieval": retrieval_metrics(mapped, y),
+ }
+
+ if args.prefix:
+ prefix, prefix_state = load_prefix(args.prefix, args.device)
+ if prefix_state["text_model"] != text["model"]:
+ raise ValueError("Prefix adapter and text feature model differ")
+ tokenizer = AutoTokenizer.from_pretrained(text["model"])
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ dtype = dtype_for_device(args.device)
+ lm = AutoModelForCausalLM.from_pretrained(
+ text["model"], torch_dtype=dtype
+ ).to(args.device)
+ lm.eval()
+ generated: list[str] = []
+ n = min(args.generation_samples, len(rows))
+ with torch.inference_mode():
+ for chunk in mapped[:n].split(16):
+ prefix_embeds = prefix(chunk.to(args.device)).to(dtype)
+ attention = torch.ones(
+ prefix_embeds.shape[:2],
+ dtype=torch.long,
+ device=args.device,
+ )
+ output = lm.generate(
+ inputs_embeds=prefix_embeds,
+ attention_mask=attention,
+ max_new_tokens=args.max_new_tokens,
+ do_sample=False,
+ eos_token_id=tokenizer.eos_token_id,
+ pad_token_id=tokenizer.pad_token_id,
+ )
+ generated.extend(tokenizer.batch_decode(output, skip_special_tokens=True))
+
+ caption_lookup = {
+ int(row): caps
+ for row, caps in zip(text["rows"], text["all_captions"])
+ }
+ records = []
+ scores = []
+ for row, caption in zip(rows[:n], generated):
+ refs = caption_lookup[int(row)]
+ score = unigram_f1(caption, refs)
+ scores.append(score)
+ records.append(
+ {
+ "row": int(row),
+ "generated": caption,
+ "references": refs,
+ "unigram_f1": score,
+ }
+ )
+ result["generation"] = {
+ "samples": n,
+ "mean_best_reference_unigram_f1": sum(scores) / max(len(scores), 1),
+ "examples": records[:25],
+ }
+
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ write_json(args.output, result)
+ print(json.dumps(result, indent=2, ensure_ascii=False))
+
+
+if __name__ == "__main__":
+ main()