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()